From f37d29dfada82d1587ef151324c2b4a37c4fdb0d Mon Sep 17 00:00:00 2001 From: Pramod Date: Thu, 16 Jul 2026 22:53:33 +0530 Subject: [PATCH 1/5] light weight scheduler and outbox sweeper (part 1) --- Makefile | 4 +- TECHNICAL_DEBT.md | 29 +++++-- docker/localstack/init-aws.sh | 6 +- frontend/web/package.json | 2 + frontend/web/pnpm-lock.yaml | 19 +++++ .../partners/api/IPartnersRepository.ts | 1 + .../src/features/partners/api/partnerHooks.ts | 7 ++ .../src/features/partners/api/partnersApi.ts | 7 ++ .../partners/components/As2PartnerDetails.tsx | 61 +++++++++---- .../components/CreatePartnerModal.tsx | 85 +++++++++++++------ .../components/CreatePartnershipModal.tsx | 18 ++-- .../components/PartnershipDetails.tsx | 16 ++-- frontend/web/src/features/partners/types.ts | 1 + .../src/features/platform/api/configHooks.ts | 34 -------- frontend/web/src/routeTree.gen.ts | 21 +++++ frontend/web/src/routes/platform.tsx | 4 +- .../42c7e50a7b1c_global_initial_schema.py | 29 +++++++ .../f966e8446341_tenant_initial_schema.py | 22 ++--- libs/database/src/database/models/__init__.py | 17 ++-- .../src/database/models/control_plane.py | 2 +- .../src/database/models/data_plane.py | 10 +-- libs/domain/src/domain/models.py | 4 +- .../src/pipeline/adapters/repository.py | 12 +-- .../src/pipeline/core/delivery/router.py | 13 ++- libs/pipeline/src/pipeline/core/saga.py | 6 +- .../pipeline/core/transformation/inbound.py | 2 +- .../pipeline/core/transformation/outbound.py | 25 ++---- .../pipeline/src/pipeline/ports/repository.py | 2 +- pyproject.toml | 4 +- .../src/api/adapters/api_token_repository.py | 2 +- services/api/src/api/adapters/http/dtos.py | 15 +++- .../api/src/api/adapters/outbox_repository.py | 51 +++++++++-- .../api/adapters/transaction_repository.py | 7 +- services/api/src/api/cdc_relay.py | 3 +- .../api/core/services/as2_partner_service.py | 8 +- .../core/services/as2_partnership_service.py | 6 +- .../core/services/inbound_route_service.py | 6 +- .../core/services/outbound_route_service.py | 6 +- .../src/api/core/services/routing_resolver.py | 11 ++- .../api/core/services/sftp_partner_service.py | 4 +- .../src/api/core/services/webhook_service.py | 6 +- services/api/src/api/core/uow.py | 16 +++- services/api/src/api/domain/models.py | 3 +- services/api/src/api/main.py | 4 + .../trading_partners/platform/__init__.py | 4 +- .../trading_partners/platform/as2_partners.py | 32 ++++++- .../trading_partners/platform/config.py | 40 --------- .../src/api/services/api_receiver_service.py | 2 +- .../src/api/services/as2_receiver_service.py | 3 +- 49 files changed, 437 insertions(+), 255 deletions(-) delete mode 100644 frontend/web/src/features/platform/api/configHooks.ts delete mode 100644 services/api/src/api/routers/trading_partners/platform/config.py diff --git a/Makefile b/Makefile index cc8cea31..f195e794 100644 --- a/Makefile +++ b/Makefile @@ -57,7 +57,7 @@ dev-web: dev-worker-orchestrator: @echo "Starting Orchestrator Worker for local development..." - ENVIRONMENT=development PYTHONPATH=services/workers/orchestrator/src:libs/database/src:libs/config/src:libs/pipeline/src:libs/domain/src:libs/transformer/src uv run python services/workers/orchestrator/src/worker/main.py + ENVIRONMENT=development PYTHONPATH=services/workers/orchestrator/src:libs/database/src:libs/config/src:libs/pipeline/src:libs/domain/src:libs/transformer/src:libs/scheduler/src uv run python services/workers/orchestrator/src/worker/main.py dev-worker-compute: @echo "Starting Compute Worker for local development..." @@ -92,7 +92,7 @@ db-sqs-reset: db-reset sqs-purge clear-data: @echo "Clearing data plane tables (edi_message, edi_json, api_gateway, outbox) and purging SQS..." - uv run python scripts/clear_data.py + uv run python scripts/clear_data.py --i-am-sure seed: db-init diff --git a/TECHNICAL_DEBT.md b/TECHNICAL_DEBT.md index aa2dda2d..2c65ee77 100644 --- a/TECHNICAL_DEBT.md +++ b/TECHNICAL_DEBT.md @@ -22,16 +22,17 @@ To achieve the long-term vision of completely merging and modernizing `bots_core ## 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. +The Outbox Sweeper has been implemented as a fallback to Debezium (CDC). It iterates over all shards to relay `PENDING` outbox events to SQS. However, the current implementation iterates over shards sequentially (`for shard in shards: await self._sweep_shard(...)`). As the number of shards grows in the multi-tenant architecture, this sequential sweep will take longer and potentially exceed the polling interval, causing lag. **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. +Refactor the Outbox Sweeper to use Bounded Concurrency or Distributed Job Fan-out: +1. **Bounded Concurrency:** Run sweeps concurrently using `asyncio.gather` bounded by an `asyncio.Semaphore` so multiple shards are swept at once without exhausting resources. +2. **Distributed Job Fan-out:** Instead of a single job, spawn a `ScheduledJob` for each shard dynamically, allowing multiple orchestrator pods to load balance the shard sweeping. +3. **Garbage Collector (Pending):** A cleanup task that runs `DELETE FROM outbox WHERE status = 'COMPLETED'` for events older than 7 days to prevent unbounded database growth is still needed. -**Estimated Effort:** Low +**Estimated Effort:** Low-Medium **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. +**Impact:** Prevents the sweeper from falling behind as the number of database shards scales, ensuring enterprise-grade multi-tenant reliability. ## 3. AS2 Protocol @@ -80,3 +81,19 @@ Currently, the bots engine does not support a lightweight validation mode (e.g., ### UnitOfWork Architecture (Control Plane vs Data Plane Naming) Currently, the `UnitOfWork` (and its underlying SQL Alchemy repositories) leak infrastructure/deployment boundaries ("Control Plane" and "Data Plane") into domain business logic. We have giant God-objects like `SqlAlchemyControlPlaneRepository` inheriting from 10+ distinct repositories, causing namespace collisions and violating SOLID principles (Single Responsibility Principle). **Future Action:** Refactor `UnitOfWork` to remove `control_plane` and `data_plane` concepts from class names and properties. Use Composition to expose distinct Bounded Contexts (e.g., `self.trading_partners`, `self.transactions`, `self.routes`) instead of lumping them into control/data plane buckets. + +## 8. Hybrid SQS Tenancy (Dynamic Queue Resolution) + +**Description:** +Currently, both inbound and outbound events are routed to a static, shared SQS queue (e.g., `TransformOrchestrationQueue`). In a multi-tenant environment, a massive batch of outbound events from one tenant can block critical inbound processing for all other tenants (the "noisy neighbor" problem). + +**Proposed Resolution:** +Implement a Hybrid SQS routing model that dynamically resolves the target queue based on the tenant's tier and the event direction: +1. **Dynamic Queue Resolver:** The CDC Relay and Outbox Sweeper should read the `tenant_id` from the outbox event and lookup the tenant's tier (cached in memory). +2. **Standard Tenants:** Route inbound events to `standard-inbound-queue` and outbound events to `standard-outbound-queue`. +3. **Enterprise Tenants:** Route events to dedicated queues (e.g., `enterprise-{tenant_name}-inbound-queue`). +4. **Dedicated Workers:** Deploy separate worker pods for standard inbound, standard outbound, and dedicated enterprise queues to provide strict compute isolation. + +**Estimated Effort:** Medium +**Estimated Time:** 3 to 5 days +**Impact:** Essential for enterprise-grade SaaS scaling. Guarantees compute isolation for enterprise customers and bulkheads heavy outbound processing from blocking high-priority inbound traffic. diff --git a/docker/localstack/init-aws.sh b/docker/localstack/init-aws.sh index 932aafef..1f4d1c2e 100755 --- a/docker/localstack/init-aws.sh +++ b/docker/localstack/init-aws.sh @@ -14,7 +14,11 @@ awslocal sqs create-queue --queue-name CDC-DLQ awslocal sqs create-queue --queue-name TransformOrchestrationQueue-DLQ TRANSFORM_DLQ_ARN=$(awslocal sqs get-queue-attributes --queue-url http://localhost:4566/000000000000/TransformOrchestrationQueue-DLQ --attribute-names QueueArn --query 'Attributes.QueueArn' --output text) -awslocal sqs create-queue --queue-name TransformOrchestrationQueue --attributes "{\"RedrivePolicy\":\"{\\\"deadLetterTargetArn\\\":\\\"$TRANSFORM_DLQ_ARN\\\",\\\"maxReceiveCount\\\":\\\"3\\\"}\"}" + awslocal sqs create-queue --queue-name TransformOrchestrationQueue --attributes "{\"RedrivePolicy\":\"{\\\"deadLetterTargetArn\\\":\\\"$TRANSFORM_DLQ_ARN\\\",\\\"maxReceiveCount\\\":\\\"3\\\"}\"}" + + awslocal sqs create-queue --queue-name TransformComputeQueue-DLQ + COMPUTE_DLQ_ARN=$(awslocal sqs get-queue-attributes --queue-url http://localhost:4566/000000000000/TransformComputeQueue-DLQ --attribute-names QueueArn --query 'Attributes.QueueArn' --output text) + awslocal sqs create-queue --queue-name TransformComputeQueue --attributes "{\"RedrivePolicy\":\"{\\\"deadLetterTargetArn\\\":\\\"$COMPUTE_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) diff --git a/frontend/web/package.json b/frontend/web/package.json index 7fc080d2..bf1a7248 100644 --- a/frontend/web/package.json +++ b/frontend/web/package.json @@ -32,6 +32,7 @@ "date-fns": "^4.4.0", "lucide-react": "^1.21.0", "next-themes": "^0.4.6", + "node-forge": "^1.4.0", "oidc-client-ts": "^3.5.0", "react": "^19.2.7", "react-dom": "^19.2.7", @@ -44,6 +45,7 @@ "devDependencies": { "@tanstack/router-plugin": "^1.168.19", "@types/node": "^24.13.2", + "@types/node-forge": "^1.3.14", "@types/react": "^19.2.17", "@types/react-dom": "^19.2.3", "@vitejs/plugin-react": "^6.0.2", diff --git a/frontend/web/pnpm-lock.yaml b/frontend/web/pnpm-lock.yaml index 113de175..f26d252a 100644 --- a/frontend/web/pnpm-lock.yaml +++ b/frontend/web/pnpm-lock.yaml @@ -74,6 +74,9 @@ importers: next-themes: specifier: ^0.4.6 version: 0.4.6(react-dom@19.2.7(react@19.2.7))(react@19.2.7) + node-forge: + specifier: ^1.4.0 + version: 1.4.0 oidc-client-ts: specifier: ^3.5.0 version: 3.5.0 @@ -105,6 +108,9 @@ importers: '@types/node': specifier: ^24.13.2 version: 24.13.2 + '@types/node-forge': + specifier: ^1.3.14 + version: 1.3.14 '@types/react': specifier: ^19.2.17 version: 19.2.17 @@ -1161,6 +1167,9 @@ packages: '@types/estree@1.0.9': resolution: {integrity: sha512-GhdPgy1el4/ImP05X05Uw4cw2/M93BCUmnEvWZNStlCzEKME4Fkk+YpoA5OiHNQmoS7Cafb8Xa3Pya8m1Qrzeg==} + '@types/node-forge@1.3.14': + resolution: {integrity: sha512-mhVF2BnD4BO+jtOp7z1CdzaK4mbuK0LLQYAvdOLqHTavxFNq4zA1EmYkpnFjP8HOUzedfQkRnp0E2ulSAYSzAw==} + '@types/node@24.13.2': resolution: {integrity: sha512-fRa09kZTgu8o71KFcDjUFuc7F+dEbZYZmkI0mg5YBTRs0yMKjYHsq/c0urDKeDb+D5qVgXOdFcuu+DZPKOITwA==} @@ -1678,6 +1687,10 @@ packages: react: ^16.8 || ^17 || ^18 || ^19 || ^19.0.0-rc react-dom: ^16.8 || ^17 || ^18 || ^19 || ^19.0.0-rc + node-forge@1.4.0: + resolution: {integrity: sha512-LarFH0+6VfriEhqMMcLX2F7SwSXeWwnEAJEsYm5QKWchiVYVvJyV9v7UDvUv+w5HO23ZpQTXDv/GxdDdMyOuoQ==} + engines: {node: '>= 6.13.0'} + node-releases@2.0.50: resolution: {integrity: sha512-J6l92tKHX6w8Jy5nO1Vuc01NoIiRGi/d6qBKVxh+IQ8Cr3b6HbVNfKiF8ZpFKufTwpwxMmce2W3iQZ861ZRyTg==} engines: {node: '>=18'} @@ -3131,6 +3144,10 @@ snapshots: '@types/estree@1.0.9': {} + '@types/node-forge@1.3.14': + dependencies: + '@types/node': 24.13.2 + '@types/node@24.13.2': dependencies: undici-types: 7.18.2 @@ -3593,6 +3610,8 @@ snapshots: react: 19.2.7 react-dom: 19.2.7(react@19.2.7) + node-forge@1.4.0: {} + node-releases@2.0.50: {} normalize-path@3.0.0: {} diff --git a/frontend/web/src/features/partners/api/IPartnersRepository.ts b/frontend/web/src/features/partners/api/IPartnersRepository.ts index 40aed236..886d8927 100644 --- a/frontend/web/src/features/partners/api/IPartnersRepository.ts +++ b/frontend/web/src/features/partners/api/IPartnersRepository.ts @@ -32,6 +32,7 @@ export interface IPartnersRepository { // Certificates exportCertificates(partnerId: string): Promise; rotateCertificates(partnerId: string, payload: RotateCertPayload): Promise; + generateCertificate(as2Id: string): Promise<{ public_cert_pem: string; private_key_vault_ref: string }>; // Tenant Partners getTenantPartners(): Promise; diff --git a/frontend/web/src/features/partners/api/partnerHooks.ts b/frontend/web/src/features/partners/api/partnerHooks.ts index 7698eb74..f988cf43 100644 --- a/frontend/web/src/features/partners/api/partnerHooks.ts +++ b/frontend/web/src/features/partners/api/partnerHooks.ts @@ -228,6 +228,13 @@ export function useRotateCertificatesMutation() { ); } +export function useGenerateCertificateMutation() { + const repo = useRepository(); + return useMutation({ + mutationFn: (as2Id: string) => repo.generateCertificate(as2Id) + }); +} + export function useTestSftpConnectionMutation() { const repo = useRepository(); return useMutation({ diff --git a/frontend/web/src/features/partners/api/partnersApi.ts b/frontend/web/src/features/partners/api/partnersApi.ts index 6700f3e0..5cdee7cb 100644 --- a/frontend/web/src/features/partners/api/partnersApi.ts +++ b/frontend/web/src/features/partners/api/partnersApi.ts @@ -124,6 +124,13 @@ class HttpPartnersRepository implements IPartnersRepository { ); } + generateCertificate(as2Id: string): Promise<{ public_cert_pem: string; private_key_vault_ref: string }> { + return this.request( + `/api/v1/platform/trading-partners/as2/certificates/generate`, + { method: 'POST', body: JSON.stringify({ as2_id: as2Id }) }, + ); + } + // ── Tenant Partners ──────────────────────── async getTenantPartners(): Promise { const data = await this.request('/api/v1/trading-partners'); diff --git a/frontend/web/src/features/partners/components/As2PartnerDetails.tsx b/frontend/web/src/features/partners/components/As2PartnerDetails.tsx index 55f70dd8..e79ad7bf 100644 --- a/frontend/web/src/features/partners/components/As2PartnerDetails.tsx +++ b/frontend/web/src/features/partners/components/As2PartnerDetails.tsx @@ -1,20 +1,21 @@ -import { useState } from 'react'; +import React, { useState } from 'react'; import { useForm, Controller } from 'react-hook-form'; import { Dialog, DialogContent, DialogHeader, DialogTitle, DialogFooter } from '@/components/ui/dialog'; import type { AS2Partner } from '../types'; import { useCertificatesExportQuery, useUpdatePlatformPartnerMutation, useRotateCertificatesMutation } from '../api/partnerHooks'; +import { pki } from 'node-forge'; import { Copy, Download, Loader2, ChevronDown, ChevronRight, CheckCircle2, Clock, ClipboardPaste } from 'lucide-react'; import { Button } from '@/components/ui/button'; import { Input } from '@/components/ui/input'; import { Label } from '@/components/ui/label'; import { useToast } from '@/hooks/use-toast'; -import { usePlatformConfig } from '@/features/platform/api/configHooks'; +import { usePlatformSettings } from '@/features/platform/api/settingsHooks'; import { Combobox } from '@/components/ui/combobox'; import { CertificateInput } from './CertificateInput'; export function As2PartnerDetails({ partner, onCancel }: { partner: AS2Partner, onCancel?: () => void }) { const { toast } = useToast(); - const { data: platformConfig } = usePlatformConfig(); + const { data: platformSettings } = usePlatformSettings(); const updatePlatform = useUpdatePlatformPartnerMutation(); const rotateCertificates = useRotateCertificatesMutation(); @@ -161,7 +162,7 @@ export function As2PartnerDetails({ partner, onCancel }: { partner: AS2Partner, control={control} render={({ field }) => ( - Status + Status Description + Issued + Expires + @@ -280,6 +284,19 @@ function CertificateRow({ const [expanded, setExpanded] = useState(false); const { toast } = useToast(); + const certInfo = React.useMemo(() => { + if (!publicPem || !publicPem.includes('-----BEGIN CERTIFICATE-----')) return null; + try { + const cert = pki.certificateFromPem(publicPem); + return { + notBefore: cert.validity.notBefore.toLocaleDateString(undefined, { year: 'numeric', month: 'short', day: 'numeric' }), + notAfter: cert.validity.notAfter.toLocaleDateString(undefined, { year: 'numeric', month: 'short', day: 'numeric' }), + }; + } catch (e) { + return null; + } + }, [publicPem]); + const handleCopy = (text: string, label: string) => { navigator.clipboard.writeText(text); toast({ title: 'Copied', description: `${label} copied to clipboard.` }); @@ -327,21 +344,35 @@ function CertificateRow({ )} - -
- {role === 'Active' ? 'Actively used for signing and decryption' : 'In grace period for legacy traffic'} -
- - {expanded ? 'Hide Details' : 'View Details'} - - {expanded ? : } -
+ + {role === 'Active' ? 'Actively used for signing and decryption' : 'In grace period for legacy traffic'} + + + {certInfo ? ( + {certInfo.notBefore} + ) : ( + Unknown + )} + + + {certInfo ? ( + {certInfo.notAfter} + ) : ( + Unknown + )} + + +
+ + {expanded ? 'Hide Details' : 'View Details'} + + {expanded ? : }
{expanded && ( - +
diff --git a/frontend/web/src/features/partners/components/CreatePartnerModal.tsx b/frontend/web/src/features/partners/components/CreatePartnerModal.tsx index 5dd8b84c..35c7a9be 100644 --- a/frontend/web/src/features/partners/components/CreatePartnerModal.tsx +++ b/frontend/web/src/features/partners/components/CreatePartnerModal.tsx @@ -3,34 +3,31 @@ import { Input } from '@/components/ui/input'; import { Label } from '@/components/ui/label'; import { FormModal } from '@/components/ui/form-modal'; import { CertificateInput } from './CertificateInput'; -import { useCreatePlatformPartnerMutation } from '../api/partnerHooks'; -import { usePlatformConfig } from '@/features/platform/api/configHooks'; +import { useCreatePlatformPartnerMutation, useGenerateCertificateMutation } from '../api/partnerHooks'; +import { usePlatformSettings } from '@/features/platform/api/settingsHooks'; import { useToast } from '@/hooks/use-toast'; import { Combobox } from '@/components/ui/combobox'; -import { useEffect } from 'react'; - +import { Button } from '@/components/ui/button'; +import { Loader2 } from 'lucide-react'; export function CreatePartnerModal({ existingAs2Ids = [] }: { existingAs2Ids?: string[] }) { const [isOpen, setIsOpen] = useState(false); const [isLocal, setIsLocal] = useState(false); const [certPem, setCertPem] = useState(''); + const [privateKeyVaultRef, setPrivateKeyVaultRef] = useState(null); const [as2Id, setAs2Id] = useState(''); const [url, setUrl] = useState(''); const isDuplicate = existingAs2Ids.includes(as2Id); - const { data: platformConfig } = usePlatformConfig(); + const { data: platformSettings } = usePlatformSettings(); const { toast } = useToast(); const createPartner = useCreatePlatformPartnerMutation(); - - useEffect(() => { - if (isLocal && !url && platformConfig?.available_as2_receive_urls?.length) { - setUrl(platformConfig.available_as2_receive_urls[0]); - } - }, [isLocal, platformConfig, url]); + const generateCert = useGenerateCertificateMutation(); const reset = () => { setIsLocal(false); setCertPem(''); + setPrivateKeyVaultRef(null); setAs2Id(''); setUrl(''); }; @@ -63,7 +60,8 @@ export function CreatePartnerModal({ existingAs2Ids = [] }: { existingAs2Ids?: s as2_id: data.get('as2_id') as string, is_local: isLocal, url: url, - public_cert_pem: isLocal ? undefined : certPem, + public_cert_pem: isLocal && privateKeyVaultRef ? certPem : isLocal ? undefined : certPem, + vault_key_ref: privateKeyVaultRef || undefined, }, { onSuccess: () => { @@ -94,8 +92,19 @@ export function CreatePartnerModal({ existingAs2Ids = [] }: { existingAs2Ids?: s role="switch" aria-checked={isLocal} onClick={() => { - setIsLocal(!isLocal); - if (!isLocal) setCertPem(''); + const nextIsLocal = !isLocal; + setIsLocal(nextIsLocal); + if (nextIsLocal) { + setCertPem(''); + setPrivateKeyVaultRef(null); + if (!url && platformSettings?.available_as2_receive_urls?.length) { + setUrl(platformSettings.available_as2_receive_urls[0]); + } + } else { + if (platformSettings?.available_as2_receive_urls?.includes(url)) { + setUrl(''); + } + } }} className={`relative inline-flex h-7 w-[90px] shrink-0 cursor-pointer items-center rounded-full border transition-colors duration-200 ease-in-out focus:outline-none focus:ring-2 focus:ring-offset-2 ${isLocal ? 'bg-indigo-50 border-indigo-200 focus:ring-indigo-200' : 'bg-violet-50 border-violet-200 focus:ring-violet-200'}`} > @@ -139,7 +148,7 @@ export function CreatePartnerModal({ existingAs2Ids = [] }: { existingAs2Ids?: s {isLocal ? ( - {isLocal ? ( -
-

- A new certificate will be automatically generated and assigned when you save this local station. -

-
- ) : ( - - )} + { + if (!as2Id.trim()) { + toast({ title: 'Error', description: 'Please enter an AS2 ID first to use as the Common Name.', variant: 'destructive' }); + return; + } + generateCert.mutate(as2Id, { + onSuccess: (res) => { + setCertPem(res.public_cert_pem); + setPrivateKeyVaultRef(res.private_key_vault_ref); + toast({ title: 'Certificate Generated', description: 'The certificate has been generated and populated.' }); + }, + onError: () => { + toast({ title: 'Error', description: 'Failed to generate certificate.', variant: 'destructive' }); + } + }); + }} + > + {generateCert.isPending ? : null} + Generate Certificate + + ) : undefined + } + />
); diff --git a/frontend/web/src/features/partners/components/CreatePartnershipModal.tsx b/frontend/web/src/features/partners/components/CreatePartnershipModal.tsx index 74e8f6e2..2173213f 100644 --- a/frontend/web/src/features/partners/components/CreatePartnershipModal.tsx +++ b/frontend/web/src/features/partners/components/CreatePartnershipModal.tsx @@ -4,7 +4,7 @@ import { Label } from '@/components/ui/label' import { FormModal } from '@/components/ui/form-modal' import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select" import { useCreatePlatformPartnershipMutation } from '../api/partnerHooks' -import { usePlatformConfig } from '@/features/platform/api/configHooks' +import { usePlatformSettings } from '@/features/platform/api/settingsHooks' import { Combobox } from '@/components/ui/combobox' import { SearchableSelect } from '@/components/ui/searchable-select' @@ -22,21 +22,21 @@ export function CreatePartnershipModal({ availablePartners }: CreatePartnershipM const [encryptionAlgorithm, setEncryptionAlgorithm] = useState('AES256') const [signatureAlgorithm, setSignatureAlgorithm] = useState('SHA256') - const { data: platformConfig } = usePlatformConfig(); + const { data: platformSettings } = usePlatformSettings(); const createPartnership = useCreatePlatformPartnershipMutation(); useEffect(() => { - if (mdnType === 'ASYNC' && !mdnUrl && platformConfig?.available_as2_receive_urls?.length) { - setMdnUrl(platformConfig.available_as2_receive_urls[0]); + if (mdnType === 'ASYNC' && !mdnUrl && platformSettings?.available_as2_receive_urls?.length) { + setMdnUrl(platformSettings.available_as2_receive_urls[0]); } - }, [platformConfig, mdnUrl, mdnType]); + }, [platformSettings, mdnUrl, mdnType]); const reset = () => { setName('') setLocalPartnerId('') setRemotePartnerId('') setMdnType('SYNC') - setMdnUrl(platformConfig?.available_as2_receive_urls?.[0] || '') + setMdnUrl(platformSettings?.available_as2_receive_urls?.[0] || '') setEncryptionAlgorithm('AES256') setSignatureAlgorithm('SHA256') } @@ -133,7 +133,7 @@ export function CreatePartnershipModal({ availablePartners }: CreatePartnershipM
- {(platformConfig?.supported_as2_encryption_algorithms || []).map(o => ( + {(platformSettings?.supported_as2_encryption_algorithms || []).map(o => ( {o.label} ))} @@ -166,7 +166,7 @@ export function CreatePartnershipModal({ availablePartners }: CreatePartnershipM - {(platformConfig?.supported_as2_signature_algorithms || []).map(o => ( + {(platformSettings?.supported_as2_signature_algorithms || []).map(o => ( {o.label} ))} diff --git a/frontend/web/src/features/partners/components/PartnershipDetails.tsx b/frontend/web/src/features/partners/components/PartnershipDetails.tsx index 97875072..1af25925 100644 --- a/frontend/web/src/features/partners/components/PartnershipDetails.tsx +++ b/frontend/web/src/features/partners/components/PartnershipDetails.tsx @@ -9,7 +9,7 @@ import { Label } from '@/components/ui/label'; import { Loader2, Zap, CheckCircle2, XCircle, Trash2 } from 'lucide-react'; import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from '@/components/ui/select'; import { useToast } from '@/hooks/use-toast'; -import { usePlatformConfig } from '@/features/platform/api/configHooks'; +import { usePlatformSettings } from '@/features/platform/api/settingsHooks'; import { Combobox } from '@/components/ui/combobox'; import { SearchableSelect } from '@/components/ui/searchable-select'; import { EdiEditorPane } from '@/components/ui/edi-editor-pane'; @@ -24,7 +24,7 @@ export function PartnershipDetails({ partnership, availablePartners, onCancel }: const { toast } = useToast(); const updatePartnership = useUpdatePlatformPartnershipMutation(); const testConnection = useTestAs2PartnershipConnectionMutation(); - const { data: platformConfig } = usePlatformConfig(); + const { data: platformSettings } = usePlatformSettings(); const isSubmitting = updatePartnership.isPending; // Custom Payload State @@ -47,10 +47,10 @@ export function PartnershipDetails({ partnership, availablePartners, onCancel }: const mdnUrl = watch('mdn_url'); useEffect(() => { - if (mdnType === 'ASYNC' && !mdnUrl && platformConfig?.available_as2_receive_urls?.length) { - setValue('mdn_url', platformConfig.available_as2_receive_urls[0], { shouldDirty: true }); + if (mdnType === 'ASYNC' && !mdnUrl && platformSettings?.available_as2_receive_urls?.length) { + setValue('mdn_url', platformSettings.available_as2_receive_urls[0], { shouldDirty: true }); } - }, [mdnType, mdnUrl, setValue, platformConfig]); + }, [mdnType, mdnUrl, setValue, platformSettings]); const onSubmit = (formData: any) => { const payload: any = {}; if (formData.name !== partnership.name) payload.name = formData.name; @@ -161,7 +161,7 @@ export function PartnershipDetails({ partnership, availablePartners, onCancel }: control={control} render={({ field }) => ( - {(platformConfig?.supported_as2_encryption_algorithms || []).map(o => ( + {(platformSettings?.supported_as2_encryption_algorithms || []).map(o => ( {o.label} ))} @@ -197,7 +197,7 @@ export function PartnershipDetails({ partnership, availablePartners, onCancel }: setLocalInterval(e.target.value)} + disabled={!localEnabled} + className="w-20 h-8" + /> + +
+ + +
+ + + + + + Job ID + Name + Status + Error Message + Next Run At + Locked By + + + + {jobs.map((job: any) => ( + + {job.id} + {job.name} + + + {job.status} + + + + {job.error_message || '-'} + + {job.next_run_at ? new Date(job.next_run_at).toLocaleString() : 'Immediate'} + {job.locked_by || '-'} + + ))} + {jobs.length === 0 && ( + + + No background jobs found. + + + )} + +
+
+ +
+ ); +}; diff --git a/frontend/web/src/routes/platform/scheduler.tsx b/frontend/web/src/routes/platform/scheduler.tsx new file mode 100644 index 00000000..630cda17 --- /dev/null +++ b/frontend/web/src/routes/platform/scheduler.tsx @@ -0,0 +1,6 @@ +import { createFileRoute } from '@tanstack/react-router'; +import { SchedulerDashboard } from '@/features/platform/components/SchedulerDashboard'; + +export const Route = createFileRoute('/platform/scheduler')({ + component: SchedulerDashboard, +}); diff --git a/libs/database/src/database/migrations/global/versions/42c7e50a7b1c_global_initial_schema.py b/libs/database/src/database/migrations/global/versions/42c7e50a7b1c_global_initial_schema.py index c9c88cc4..3095a393 100644 --- a/libs/database/src/database/migrations/global/versions/42c7e50a7b1c_global_initial_schema.py +++ b/libs/database/src/database/migrations/global/versions/42c7e50a7b1c_global_initial_schema.py @@ -374,6 +374,11 @@ def upgrade() -> None: def downgrade() -> None: """Downgrade schema.""" # ### commands auto generated by Alembic - please adjust! ### + op.drop_index(op.f("ix_scheduled_jobs_status"), table_name="scheduled_jobs") + op.drop_index(op.f("ix_scheduled_jobs_next_run_at"), table_name="scheduled_jobs") + op.drop_index(op.f("ix_scheduled_jobs_name"), table_name="scheduled_jobs") + op.drop_table("scheduled_jobs") + op.drop_table("platform_settings") op.drop_index( "ix_outbound_routes_unique_trading_partner_id", table_name="outbound_routes", diff --git a/libs/database/src/database/models/platform_settings.py b/libs/database/src/database/models/platform_settings.py new file mode 100644 index 00000000..1f568077 --- /dev/null +++ b/libs/database/src/database/models/platform_settings.py @@ -0,0 +1,16 @@ +from typing import Any + +from sqlalchemy import JSON, String +from sqlalchemy.orm import Mapped, mapped_column + +from database.models.common import TimestampMixin +from database.models.control_plane import GlobalBase + + +class PlatformSettings(GlobalBase, TimestampMixin): + __tablename__ = "platform_settings" + + key: Mapped[str] = mapped_column(String, primary_key=True) + value: Mapped[dict[str, Any] | list[Any] | str | int | bool | None] = mapped_column( + JSON, nullable=True + ) diff --git a/libs/database/src/database/models/scheduled_job.py b/libs/database/src/database/models/scheduled_job.py new file mode 100644 index 00000000..48f1c74e --- /dev/null +++ b/libs/database/src/database/models/scheduled_job.py @@ -0,0 +1,32 @@ +import uuid +from datetime import datetime +from typing import Any + +from sqlalchemy import JSON, DateTime, Integer, String +from sqlalchemy.orm import Mapped, mapped_column + +from database.models.control_plane import GlobalBase + + +class ScheduledJob(GlobalBase): + __tablename__ = "scheduled_jobs" + + id: Mapped[uuid.UUID] = mapped_column(primary_key=True, default=uuid.uuid4) + name: Mapped[str] = mapped_column(String, index=True) + payload: Mapped[dict[str, Any]] = mapped_column(JSON, default=dict) + status: Mapped[str] = mapped_column(String, index=True, default="PENDING") + + next_run_at: Mapped[datetime | None] = mapped_column( + DateTime(timezone=True), nullable=True, index=True + ) + locked_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True) + locked_by: Mapped[str | None] = mapped_column(String, nullable=True) + + retry_count: Mapped[int] = mapped_column(Integer, default=0) + max_retries: Mapped[int] = mapped_column(Integer, default=3) + error_message: Mapped[str | None] = mapped_column(String, nullable=True) + + created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=datetime.utcnow) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), default=datetime.utcnow, onupdate=datetime.utcnow + ) diff --git a/libs/pipeline/src/pipeline/core/transformation/outbound.py b/libs/pipeline/src/pipeline/core/transformation/outbound.py index 8ebb35ac..37d67570 100644 --- a/libs/pipeline/src/pipeline/core/transformation/outbound.py +++ b/libs/pipeline/src/pipeline/core/transformation/outbound.py @@ -41,6 +41,8 @@ async def transform(self, trace_id: str) -> None: if not trading_partner_id: trading_partner_id = routing_meta.get("trading_partner_id") + route_config = None + outbound_route = None if trading_partner_id: route_config = await self.repository.get_outbound_edi_header_by_route_or_partner( trading_partner_id=trading_partner_id, tenant_id=tenant_id diff --git a/libs/scheduler/README.md b/libs/scheduler/README.md new file mode 100644 index 00000000..e69de29b diff --git a/libs/scheduler/pyproject.toml b/libs/scheduler/pyproject.toml new file mode 100644 index 00000000..bf0efc34 --- /dev/null +++ b/libs/scheduler/pyproject.toml @@ -0,0 +1,23 @@ +[project] +name = "scheduler" +version = "0.1.0" +description = "Generic scheduler module for background tasks" +readme = "README.md" +requires-python = ">=3.11" +dependencies = [ + "pydantic>=2.0.0", + "SQLAlchemy>=2.0.0", + # Local dependencies + "database", +] + +[tool.uv.sources] +database = { workspace = true } + +[build-system] +requires = ["hatchling"] +build-backend = "hatchling.build" + +[tool.pytest.ini_options] +asyncio_mode = "auto" +testpaths = ["tests"] diff --git a/libs/scheduler/src/scheduler/__init__.py b/libs/scheduler/src/scheduler/__init__.py new file mode 100644 index 00000000..484b1f9f --- /dev/null +++ b/libs/scheduler/src/scheduler/__init__.py @@ -0,0 +1,14 @@ +from scheduler.adapters.repository import SqlAlchemyJobRepository +from scheduler.core.service import SchedulerWorkerService +from scheduler.domain.models import Job, JobStatus +from scheduler.ports.handler import JobHandlerPort +from scheduler.ports.repository import JobRepositoryPort + +__all__ = [ + "SchedulerWorkerService", + "Job", + "JobStatus", + "JobHandlerPort", + "JobRepositoryPort", + "SqlAlchemyJobRepository", +] diff --git a/libs/scheduler/src/scheduler/adapters/repository.py b/libs/scheduler/src/scheduler/adapters/repository.py new file mode 100644 index 00000000..e21297a0 --- /dev/null +++ b/libs/scheduler/src/scheduler/adapters/repository.py @@ -0,0 +1,121 @@ +import datetime +import logging +import uuid +from typing import Any + +from database.models.scheduled_job import ScheduledJob +from scheduler.domain.models import Job, JobStatus +from scheduler.ports.repository import JobRepositoryPort +from sqlalchemy import select, update +from sqlalchemy.ext.asyncio import AsyncEngine, async_sessionmaker + +logger = logging.getLogger(__name__) + + +class SqlAlchemyJobRepository(JobRepositoryPort): + def __init__(self, engine: AsyncEngine): + self.session_factory = async_sessionmaker(engine, expire_on_commit=False) + + def _to_domain(self, record: ScheduledJob) -> Job: + return Job( + id=record.id, + name=record.name, + payload=record.payload, + status=JobStatus(record.status), + next_run_at=record.next_run_at, + retry_count=record.retry_count, + max_retries=record.max_retries, + locked_at=record.locked_at, + locked_by=record.locked_by, + created_at=record.created_at, + updated_at=record.updated_at, + ) + + async def claim_next_job(self, worker_id: str) -> Job | None: + """ + Uses SKIP LOCKED to safely claim the next PENDING or ready job. + """ + now = datetime.datetime.now(datetime.UTC) + + async with self.session_factory() as session, session.begin(): + # Find the next + stmt = ( + select(ScheduledJob) + .where( + (ScheduledJob.status == JobStatus.PENDING.value) + & (ScheduledJob.next_run_at.is_(None) | (ScheduledJob.next_run_at <= now)) + ) + .order_by(ScheduledJob.created_at.asc()) + .limit(1) + .with_for_update(skip_locked=True) + ) + + result = await session.execute(stmt) + record = result.scalar_one_or_none() + + if not record: + return None + + # Claim it + record.status = JobStatus.RUNNING.value + record.locked_at = now + record.locked_by = worker_id + + await session.flush() + return self._to_domain(record) + + async def mark_completed(self, job_id: uuid.UUID) -> None: + async with self.session_factory() as session, session.begin(): + stmt = ( + update(ScheduledJob) + .where(ScheduledJob.id == job_id) + .values( + status=JobStatus.COMPLETED.value, + locked_at=None, + locked_by=None, + ) + ) + await session.execute(stmt) + + async def mark_failed(self, job_id: uuid.UUID, error: str) -> None: + async with self.session_factory() as session, session.begin(): + stmt = ( + update(ScheduledJob) + .where(ScheduledJob.id == job_id) + .values( + status=JobStatus.FAILED.value, + locked_at=None, + locked_by=None, + error_message=error, + ) + ) + await session.execute(stmt) + + async def schedule_job( + self, name: str, payload: dict[str, Any], next_run_at: datetime.datetime | None = None + ) -> Job: + async with self.session_factory() as session, session.begin(): + record = ScheduledJob( + name=name, + payload=payload, + status=JobStatus.PENDING.value, + next_run_at=next_run_at, + ) + session.add(record) + await session.flush() + return self._to_domain(record) + + async def reschedule(self, job_id: uuid.UUID, next_run_at: datetime.datetime) -> None: + async with self.session_factory() as session, session.begin(): + stmt = ( + update(ScheduledJob) + .where(ScheduledJob.id == job_id) + .values( + status=JobStatus.PENDING.value, + next_run_at=next_run_at, + locked_at=None, + locked_by=None, + retry_count=0, + ) + ) + await session.execute(stmt) diff --git a/libs/scheduler/src/scheduler/core/service.py b/libs/scheduler/src/scheduler/core/service.py new file mode 100644 index 00000000..4dd69ac1 --- /dev/null +++ b/libs/scheduler/src/scheduler/core/service.py @@ -0,0 +1,67 @@ +import asyncio +import contextlib +import logging + +from scheduler.ports.handler import JobHandlerPort +from scheduler.ports.repository import JobRepositoryPort + +logger = logging.getLogger(__name__) + + +class SchedulerWorkerService: + def __init__(self, repository: JobRepositoryPort, worker_id: str): + self.repository = repository + self.worker_id = worker_id + self.handlers: dict[str, JobHandlerPort] = {} + self._is_running = False + self._task: asyncio.Task[None] | None = None + + def register_handler(self, job_name: str, handler: JobHandlerPort) -> None: + self.handlers[job_name] = handler + + async def start(self, poll_interval_seconds: float = 5.0) -> None: + self._is_running = True + logger.info(f"Starting scheduler worker {self.worker_id}") + self._task = asyncio.create_task(self._poll_loop(poll_interval_seconds)) + + async def stop(self) -> None: + self._is_running = False + if self._task: + self._task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await self._task + logger.info(f"Stopped scheduler worker {self.worker_id}") + + async def _poll_loop(self, poll_interval_seconds: float) -> None: + while self._is_running: + try: + job = await self.repository.claim_next_job(worker_id=self.worker_id) + if job: + handler = self.handlers.get(job.name) + if not handler: + error_msg = f"No handler registered for job {job.name}" + logger.error(error_msg) + await self.repository.mark_failed(job.id, error=error_msg) + continue + + try: + logger.info(f"Executing job {job.name} ({job.id})") + next_run_at = await handler.execute(job) + if next_run_at: + await self.repository.reschedule(job.id, next_run_at) + logger.info( + f"Successfully rescheduled job {job.name} ({job.id}) for {next_run_at}" + ) + else: + await self.repository.mark_completed(job.id) + logger.info(f"Successfully completed job {job.name} ({job.id})") + except Exception as e: + logger.exception(f"Job {job.name} ({job.id}) failed: {e}") + await self.repository.mark_failed(job.id, error=str(e)) + else: + await asyncio.sleep(poll_interval_seconds) + except asyncio.CancelledError: + break + except Exception as e: + logger.error(f"Error in scheduler poll loop: {e}") + await asyncio.sleep(poll_interval_seconds) diff --git a/libs/scheduler/src/scheduler/domain/models.py b/libs/scheduler/src/scheduler/domain/models.py new file mode 100644 index 00000000..b256ea51 --- /dev/null +++ b/libs/scheduler/src/scheduler/domain/models.py @@ -0,0 +1,27 @@ +import uuid +from dataclasses import dataclass, field +from datetime import datetime +from enum import StrEnum +from typing import Any + + +class JobStatus(StrEnum): + PENDING = "PENDING" + RUNNING = "RUNNING" + COMPLETED = "COMPLETED" + FAILED = "FAILED" + + +@dataclass +class Job: + name: str + payload: dict[str, Any] + status: JobStatus = JobStatus.PENDING + next_run_at: datetime | None = None + retry_count: int = 0 + max_retries: int = 3 + locked_at: datetime | None = None + locked_by: str | None = None + id: uuid.UUID = field(default_factory=uuid.uuid4) + created_at: datetime | None = None + updated_at: datetime | None = None diff --git a/libs/scheduler/src/scheduler/ports/handler.py b/libs/scheduler/src/scheduler/ports/handler.py new file mode 100644 index 00000000..8642653c --- /dev/null +++ b/libs/scheduler/src/scheduler/ports/handler.py @@ -0,0 +1,11 @@ +import abc +import datetime + +from scheduler.domain.models import Job + + +class JobHandlerPort(abc.ABC): + @abc.abstractmethod + async def execute(self, job: Job) -> datetime.datetime | None: + """Execute the job. Return a datetime to reschedule it, or None to complete it.""" + pass diff --git a/libs/scheduler/src/scheduler/ports/repository.py b/libs/scheduler/src/scheduler/ports/repository.py new file mode 100644 index 00000000..a4647986 --- /dev/null +++ b/libs/scheduler/src/scheduler/ports/repository.py @@ -0,0 +1,30 @@ +import abc +import uuid +from datetime import datetime +from typing import Any + +from scheduler.domain.models import Job + + +class JobRepositoryPort(abc.ABC): + @abc.abstractmethod + async def claim_next_job(self, worker_id: str) -> Job | None: + pass + + @abc.abstractmethod + async def mark_completed(self, job_id: uuid.UUID) -> None: + pass + + @abc.abstractmethod + async def mark_failed(self, job_id: uuid.UUID, error: str) -> None: + pass + + @abc.abstractmethod + async def schedule_job( + self, name: str, payload: dict[str, Any], next_run_at: datetime | None = None + ) -> Job: + pass + + @abc.abstractmethod + async def reschedule(self, job_id: uuid.UUID, next_run_at: datetime) -> None: + pass diff --git a/services/api/src/api/adapters/outbox_repository.py b/services/api/src/api/adapters/outbox_repository.py index bdcd9bd0..a8ede093 100644 --- a/services/api/src/api/adapters/outbox_repository.py +++ b/services/api/src/api/adapters/outbox_repository.py @@ -7,16 +7,9 @@ from database.models.control_plane import ControlPlaneOutbox -class SqlAlchemyControlPlaneOutboxRepository(GlobalSqlAlchemyRepository, OutboxRepositoryPort): - """ - Outbox repository for the Control Plane (Global DB). - Writes provisioning events (AS2_PARTNER_CREATED, etc.) that are later - polled by the Provisioning Worker to replicate config to tenant shards. - """ - - def __init__(self, session: Any, model_class: Any = ControlPlaneOutbox) -> None: - super().__init__(session) - self.model_class = model_class +class SqlAlchemyOutboxRepositoryMixin: + session: Any + model_class: Any async def publish_outbox_event( self, @@ -39,7 +32,23 @@ async def publish_outbox_event( return event_id -class SqlAlchemyDataPlaneOutboxRepository(TenantSqlAlchemyRepository, OutboxRepositoryPort): +class SqlAlchemyControlPlaneOutboxRepository( + SqlAlchemyOutboxRepositoryMixin, GlobalSqlAlchemyRepository, OutboxRepositoryPort +): + """ + Outbox repository for the Control Plane (Global DB). + Writes provisioning events (AS2_PARTNER_CREATED, etc.) that are later + polled by the Provisioning Worker to replicate config to tenant shards. + """ + + def __init__(self, session: Any, model_class: Any = ControlPlaneOutbox) -> None: + super().__init__(session) + self.model_class = model_class + + +class SqlAlchemyDataPlaneOutboxRepository( + SqlAlchemyOutboxRepositoryMixin, TenantSqlAlchemyRepository, OutboxRepositoryPort +): """ Outbox repository for the Data Plane (Tenant Shard). Writes pipeline events (TRANSFORM_EVENT, DELIVER_EVENT, etc.) consumed @@ -51,23 +60,3 @@ def __init__(self, session: Any) -> None: super().__init__(session) self.model_class = DataPlaneOutbox - - async def publish_outbox_event( - self, - tenant_id: int, - event_type: str, - payload: dict[str, Any], - idempotency_key: UUID | None = None, - ) -> UUID: - event_id = uuid.uuid4() - record = self.model_class( - id=event_id, - tenant_id=tenant_id, - idempotency_key=idempotency_key or uuid.uuid4(), - event_type=event_type, - payload=payload, - status="PENDING", - ) - self.session.add(record) - await self.session.flush() - return event_id diff --git a/services/api/src/api/adapters/platform_settings_repository.py b/services/api/src/api/adapters/platform_settings_repository.py new file mode 100644 index 00000000..4ae83f62 --- /dev/null +++ b/services/api/src/api/adapters/platform_settings_repository.py @@ -0,0 +1,30 @@ +from typing import Any + +from api.ports.platform_settings_repository import PlatformSettingsRepositoryPort +from database.models.platform_settings import PlatformSettings +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + + +class SqlAlchemyPlatformSettingsRepository(PlatformSettingsRepositoryPort): + def __init__(self, global_session: AsyncSession): + self.global_session = global_session + + async def get_config(self, key: str) -> Any | None: + stmt = select(PlatformSettings).where(PlatformSettings.key == key) + result = await self.global_session.execute(stmt) + record = result.scalar_one_or_none() + if record: + return record.value + return None + + async def set_config(self, key: str, value: Any) -> None: + stmt = select(PlatformSettings).where(PlatformSettings.key == key) + result = await self.global_session.execute(stmt) + record = result.scalar_one_or_none() + if record: + record.value = value + else: + record = PlatformSettings(key=key, value=value) + self.global_session.add(record) + await self.global_session.flush() diff --git a/services/api/src/api/adapters/transaction_repository.py b/services/api/src/api/adapters/transaction_repository.py index 843fcbcd..8f548df2 100644 --- a/services/api/src/api/adapters/transaction_repository.py +++ b/services/api/src/api/adapters/transaction_repository.py @@ -129,12 +129,16 @@ def _apply_dynamic_filters(self, stmt: Any, model: Any, filters: list[dict[str, if hasattr(model, "sender_id") and hasattr(model, "receiver_id"): has_gs = hasattr(model, "gs_sender_id") and hasattr(model, "gs_receiver_id") + has_tp = hasattr(model, "trading_partner_id") + if operator == "eq": conds = [model.sender_id == value, model.receiver_id == value] if has_gs: conds.extend( [model.gs_sender_id == value, model.gs_receiver_id == value] ) + if has_tp: + conds.append(model.trading_partner_id == value) stmt = stmt.where(or_(*conds)) elif operator == "neq": @@ -143,6 +147,8 @@ def _apply_dynamic_filters(self, stmt: Any, model: Any, filters: list[dict[str, conds.extend( [model.gs_sender_id != value, model.gs_receiver_id != value] ) + if has_tp: + conds.append(model.trading_partner_id != value) stmt = stmt.where(and_(*conds)) elif operator == "contains": @@ -161,6 +167,8 @@ def _apply_dynamic_filters(self, stmt: Any, model: Any, filters: list[dict[str, model.gs_receiver_id.ilike(pattern, escape="\\"), ] ) + if has_tp: + conds.append(model.trading_partner_id.ilike(pattern, escape="\\")) stmt = stmt.where(or_(*conds)) elif operator == "in" and isinstance(value, list): @@ -169,6 +177,8 @@ def _apply_dynamic_filters(self, stmt: Any, model: Any, filters: list[dict[str, conds.extend( [model.gs_sender_id.in_(value), model.gs_receiver_id.in_(value)] ) + if has_tp: + conds.append(model.trading_partner_id.in_(value)) stmt = stmt.where(or_(*conds)) continue diff --git a/services/api/src/api/ports/platform_settings_repository.py b/services/api/src/api/ports/platform_settings_repository.py new file mode 100644 index 00000000..a1e319ef --- /dev/null +++ b/services/api/src/api/ports/platform_settings_repository.py @@ -0,0 +1,12 @@ +import abc +from typing import Any + + +class PlatformSettingsRepositoryPort(abc.ABC): + @abc.abstractmethod + async def get_config(self, key: str) -> Any | None: + pass + + @abc.abstractmethod + async def set_config(self, key: str, value: Any) -> None: + pass diff --git a/services/api/src/api/routers/platform/__init__.py b/services/api/src/api/routers/platform/__init__.py new file mode 100644 index 00000000..703a3308 --- /dev/null +++ b/services/api/src/api/routers/platform/__init__.py @@ -0,0 +1,8 @@ +from fastapi import APIRouter + +from .scheduler import router as scheduler_router + +router = APIRouter(prefix="/api/v1/platform") +router.include_router(scheduler_router) + +__all__ = ["router"] diff --git a/services/api/src/api/routers/platform/scheduler.py b/services/api/src/api/routers/platform/scheduler.py new file mode 100644 index 00000000..f69cab94 --- /dev/null +++ b/services/api/src/api/routers/platform/scheduler.py @@ -0,0 +1,127 @@ +import uuid +from datetime import datetime +from typing import Any + +from database.models.scheduled_job import ScheduledJob +from fastapi import APIRouter, Depends +from pydantic import BaseModel, ConfigDict +from sqlalchemy import select + +from api.core.uow import UnitOfWork +from api.dependencies import get_uow + +router = APIRouter(prefix="/scheduler", tags=["Platform Scheduler"]) + + +class JobResponse(BaseModel): + id: uuid.UUID + name: str + status: str + next_run_at: datetime | None + locked_at: datetime | None + locked_by: str | None + retry_count: int + error_message: str | None + created_at: datetime + updated_at: datetime + + model_config = ConfigDict(from_attributes=True) + + +class ConfigUpdateRequest(BaseModel): + value: dict[str, Any] | list[Any] | str | int | bool | None + + +class ConfigResponse(BaseModel): + key: str + value: Any + + +@router.get("/jobs", response_model=list[JobResponse]) +async def list_jobs(uow: UnitOfWork = Depends(get_uow)) -> list[JobResponse]: + """List all scheduled background jobs (Admin Only).""" + async with uow: + stmt = select(ScheduledJob).order_by(ScheduledJob.created_at.desc()) + result = await uow.global_session.execute(stmt) + jobs = result.scalars().all() + return [JobResponse.model_validate(job) for job in jobs] + + +@router.get("/config", response_model=list[ConfigResponse]) +async def get_all_config(uow: UnitOfWork = Depends(get_uow)) -> list[ConfigResponse]: + """Get all platform configuration values.""" + from database.models.platform_settings import PlatformSettings + + async with uow: + stmt = select(PlatformSettings).order_by(PlatformSettings.key) + result = await uow.global_session.execute(stmt) + configs = result.scalars().all() + return [ConfigResponse(key=c.key, value=c.value) for c in configs] + + +@router.get("/config/{key}", response_model=ConfigResponse) +async def get_config(key: str, uow: UnitOfWork = Depends(get_uow)) -> ConfigResponse: + """Get a specific platform configuration value.""" + async with uow: + val = await uow.platform_settings.get_config(key) + return ConfigResponse(key=key, value=val) + + +@router.put("/config/{key}", response_model=ConfigResponse) +async def update_config( + key: str, request: ConfigUpdateRequest, uow: UnitOfWork = Depends(get_uow) +) -> ConfigResponse: + """Update a platform configuration value.""" + import datetime + import uuid + + from database.models.scheduled_job import ScheduledJob + + async with uow: + await uow.platform_settings.set_config(key, request.value) + + # Event-driven scheduler integration + if key == "outbox_sweeper_enabled": + stmt = select(ScheduledJob).where(ScheduledJob.name == "outbox_sweeper") + result = await uow.global_session.execute(stmt) + job = result.scalar_one_or_none() + now = datetime.datetime.now(datetime.UTC) + + if request.value is True: + if job: + job.next_run_at = now + job.status = "PENDING" + else: + # Get interval if exists + interval_cfg = await uow.platform_settings.get_config( + "outbox_sweeper_interval_seconds" + ) + interval = interval_cfg if interval_cfg is not None else 60 + + new_job = ScheduledJob( + id=uuid.uuid4(), + name="outbox_sweeper", + payload={"interval_seconds": int(interval)}, + status="PENDING", + next_run_at=now, + retry_count=0, + max_retries=3, + created_at=now, + updated_at=now, + ) + uow.global_session.add(new_job) + else: + if job: + await uow.global_session.delete(job) + + elif key == "outbox_sweeper_interval_seconds": + stmt = select(ScheduledJob).where(ScheduledJob.name == "outbox_sweeper") + result = await uow.global_session.execute(stmt) + job = result.scalar_one_or_none() + if job: + payload = dict(job.payload) if job.payload else {} + payload["interval_seconds"] = int(str(request.value)) + job.payload = payload + + await uow.commit() + return ConfigResponse(key=key, value=request.value) 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 e2ad7449..bf4ca5ca 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 @@ -53,6 +53,18 @@ async def generate_certificate( ) +@router.delete( + "/as2/certificates/secret", + status_code=status.HTTP_204_NO_CONTENT, +) +async def delete_certificate_secret( + vault_ref: str, + vault: VaultPort = Depends(get_vault), +) -> None: + """Deletes an orphaned private key from Vault if the UI discards it before saving.""" + vault.delete_secret(vault_ref) + + @router.post( "/as2/trading-partners", response_model=AS2TradingPartnerResponse, @@ -72,7 +84,9 @@ async def create_platform_as2_partner( public_cert_pem = request.public_cert_pem private_key_vault_ref = request.private_key_vault_ref + auto_generated = False if request.is_local and not private_key_vault_ref: + auto_generated = True # Auto-generate self-signed cert if not already generated and provided private_key_bytes, public_cert_bytes = generate_self_signed_cert( common_name=request.as2_id @@ -114,7 +128,7 @@ async def create_platform_as2_partner( active=p.active, ) except IntegrityError as e: - if request.is_local and private_key_vault_ref: + if auto_generated and private_key_vault_ref: vault.delete_secret(private_key_vault_ref) raise HTTPException(status_code=400, detail="AS2 ID already exists for this tenant.") from e diff --git a/services/api/src/api/routers/trading_partners/platform/settings.py b/services/api/src/api/routers/trading_partners/platform/settings.py new file mode 100644 index 00000000..b64e1f0c --- /dev/null +++ b/services/api/src/api/routers/trading_partners/platform/settings.py @@ -0,0 +1,40 @@ +from typing import Any + +from config.settings import get_settings +from fastapi import APIRouter +from pydantic import BaseModel + +router = APIRouter(tags=["Platform Settings"]) + + +class SupportedAlgorithm(BaseModel): + value: str + label: str + + +class PlatformSettingsResponse(BaseModel): + available_as2_receive_urls: list[str] + supported_as2_encryption_algorithms: list[SupportedAlgorithm] + supported_as2_signature_algorithms: list[SupportedAlgorithm] + + +@router.get("/config", response_model=PlatformSettingsResponse) +async def get_platform_settings() -> Any: + settings = get_settings() + + # We strip trailing slashes to ensure consistent path appending + base_url = settings.server.external_url.rstrip("/") + + return PlatformSettingsResponse( + available_as2_receive_urls=[f"{base_url}/api/v1/as2/receive"], + supported_as2_encryption_algorithms=[ + SupportedAlgorithm(value="AES256", label="AES-256-CBC"), + SupportedAlgorithm(value="AES128", label="AES-128-CBC"), + SupportedAlgorithm(value="3DES", label="3DES (Legacy)"), + ], + supported_as2_signature_algorithms=[ + SupportedAlgorithm(value="SHA256", label="SHA-256"), + SupportedAlgorithm(value="SHA1", label="SHA-1 (Legacy)"), + SupportedAlgorithm(value="MD5", label="MD5 (Legacy)"), + ], + ) diff --git a/services/api/src/api/services/as2_receiver_service.py b/services/api/src/api/services/as2_receiver_service.py index fbeb8ec2..f52cf065 100644 --- a/services/api/src/api/services/as2_receiver_service.py +++ b/services/api/src/api/services/as2_receiver_service.py @@ -374,9 +374,11 @@ async def _save_transaction( # type: ignore async_gen_tenant = self.db_router.get_tenant_session(true_tenant_id, shard.name, shard.dsn) tenant_session = await anext(async_gen_tenant) try: + from api.adapters.outbox_repository import SqlAlchemyDataPlaneOutboxRepository from api.adapters.transaction_repository import SqlAlchemyTransactionRepository dp_repo = SqlAlchemyTransactionRepository(tenant_session) + outbox_repo = SqlAlchemyDataPlaneOutboxRepository(tenant_session) msg_id = await dp_repo.create_edi_message(tenant_id=true_tenant_id, payload=edi_record) outbox_payload = { @@ -386,7 +388,7 @@ async def _save_transaction( # type: ignore "receiver_id": isa_receiver, "status": "RECEIVED", } - await dp_repo.publish_outbox_event( + await outbox_repo.publish_outbox_event( tenant_id=true_tenant_id, event_type=PipelineEventType.TRANSFORM_EVENT, payload=outbox_payload, diff --git a/services/api/tests/api_fakes.py b/services/api/tests/api_fakes.py index 86ce0850..278cf761 100644 --- a/services/api/tests/api_fakes.py +++ b/services/api/tests/api_fakes.py @@ -498,7 +498,7 @@ def __init__(self): self.as2_partnerships = repo self.inbound_routes = repo self.outbound_routes = repo - self.outbox = repo + self.control_plane_outbox = repo self.sftp_partners = repo self.tenants = repo self.webhooks = repo @@ -518,5 +518,9 @@ async def __aexit__(self, exc_type, exc_val, exc_tb): async def commit(self): pass + @property + def data_plane_outbox(self): + return self.transactions + async def rollback(self): pass diff --git a/services/api/tests/test_api_receiver_service.py b/services/api/tests/test_api_receiver_service.py index 8ca6dcbe..a7876602 100644 --- a/services/api/tests/test_api_receiver_service.py +++ b/services/api/tests/test_api_receiver_service.py @@ -17,9 +17,9 @@ async def test_process_api_edi_json_success(): assert trace_id is not None mock_uow.transactions.create_edi_json.assert_awaited_once() - mock_uow.outbox.publish_outbox_event.assert_awaited_once() + mock_uow.data_plane_outbox.publish_outbox_event.assert_awaited_once() - args, kwargs = mock_uow.outbox.publish_outbox_event.call_args + args, kwargs = mock_uow.data_plane_outbox.publish_outbox_event.call_args from domain.events import PipelineEventType assert kwargs["event_type"] == PipelineEventType.TRANSFORM_EVENT diff --git a/services/api/tests/test_as2_partner_service.py b/services/api/tests/test_as2_partner_service.py index 04fa7950..f6676b63 100644 --- a/services/api/tests/test_as2_partner_service.py +++ b/services/api/tests/test_as2_partner_service.py @@ -9,7 +9,7 @@ def make_mock_uow(mock_repo: AsyncMock) -> MagicMock: uow = MagicMock() uow.as2_partners = mock_repo - uow.outbox = mock_repo + uow.control_plane_outbox = mock_repo uow.global_session = mock_repo return uow diff --git a/services/api/tests/test_as2_receiver_service.py b/services/api/tests/test_as2_receiver_service.py index a661c5a4..682c32ca 100644 --- a/services/api/tests/test_as2_receiver_service.py +++ b/services/api/tests/test_as2_receiver_service.py @@ -212,13 +212,21 @@ async def mock_async_gen(): service.db_router.get_tenant_session = MagicMock(return_value=mock_async_gen()) - with patch( - "api.adapters.transaction_repository.SqlAlchemyTransactionRepository" - ) as mock_repo_cls: + with ( + patch( + "api.adapters.transaction_repository.SqlAlchemyTransactionRepository" + ) as mock_repo_cls, + patch( + "api.adapters.outbox_repository.SqlAlchemyDataPlaneOutboxRepository" + ) as mock_outbox_cls, + ): mock_repo = AsyncMock() mock_repo.create_edi_message.return_value = "msg-1" mock_repo_cls.return_value = mock_repo + mock_outbox = AsyncMock() + mock_outbox_cls.return_value = mock_outbox + mock_partnership = MagicMock( tenant_id=1, mdn_type="SYNC", @@ -234,7 +242,7 @@ async def mock_async_gen(): ) assert res == "msg-1" mock_repo.create_edi_message.assert_awaited_once() - mock_repo.publish_outbox_event.assert_awaited_once() - args, kwargs = mock_repo.publish_outbox_event.call_args + mock_outbox.publish_outbox_event.assert_awaited_once() + args, kwargs = mock_outbox.publish_outbox_event.call_args assert kwargs["idempotency_key"] == "msg-1" mock_session.commit.assert_awaited_once() diff --git a/services/api/tests/test_inbound_flow_e2e.py b/services/api/tests/test_inbound_flow_e2e.py index cc9eb27d..a410b79a 100644 --- a/services/api/tests/test_inbound_flow_e2e.py +++ b/services/api/tests/test_inbound_flow_e2e.py @@ -192,6 +192,7 @@ async def webhook_handler(request: web.Request) -> web.Response: from unittest.mock import MagicMock from database.models.data_plane import EdiMessage + from domain.events import PipelineEventType from pipeline.adapters.http import HttpxDeliveryAdapter from pipeline.adapters.repository import SqlAlchemyRepositoryAdapter from pipeline.adapters.storage import S3StorageAdapter @@ -209,7 +210,7 @@ async def webhook_handler(request: web.Request) -> web.Response: trace_id = str(edi_msg.trace_id) try: - await translate_svc.translate(trace_id, event_type="edi_message.received") + await translate_svc.translate(trace_id, event_type=PipelineEventType.TRANSFORM_EVENT) except Exception as e: # If bots is not running, we mock it for the test if "Connection" in str(e): @@ -243,7 +244,9 @@ async def translate_json_to_edi( return b"" translate_svc.transformer = MockTransformer() - await translate_svc.translate(trace_id, event_type="edi_message.received") + await translate_svc.translate( + trace_id, event_type=PipelineEventType.TRANSFORM_EVENT + ) # 2. Manually run Deliver deliver_svc = DeliveryService( diff --git a/services/api/tests/test_provisioning_core.py b/services/api/tests/test_provisioning_core.py index 38a37130..109e0a46 100644 --- a/services/api/tests/test_provisioning_core.py +++ b/services/api/tests/test_provisioning_core.py @@ -35,7 +35,7 @@ def mock_uow(global_repo): uow.as2_partnerships = global_repo uow.inbound_routes = global_repo uow.outbound_routes = global_repo - uow.outbox = global_repo + uow.control_plane_outbox = global_repo uow.sftp_partners = global_repo uow.tenants = global_repo uow.webhooks = global_repo diff --git a/services/api/tests/test_routers_partners.py b/services/api/tests/test_routers_partners.py index 9917dd52..67c58cff 100644 --- a/services/api/tests/test_routers_partners.py +++ b/services/api/tests/test_routers_partners.py @@ -84,7 +84,7 @@ def test_list_platform_as2_partners(client, fake_uow): assert export_resp.status_code in (200, 403, 404, 501) -def test_get_platform_config(client, fake_uow): +def test_get_platform_settings(client, fake_uow): response = client.get("/api/v1/platform/trading-partners/config") assert response.status_code == 200 diff --git a/services/api/tests/test_routers_transactions.py b/services/api/tests/test_routers_transactions.py index 9f5674b5..e0e73397 100644 --- a/services/api/tests/test_routers_transactions.py +++ b/services/api/tests/test_routers_transactions.py @@ -30,7 +30,7 @@ def _make_mock_msg() -> MagicMock: m.status = "SUCCESS" m.edi_data = "TEST" m.created_at = datetime.now(UTC) - m.outbound_route_id = uuid.uuid4() + m.trading_partner_id = "TEST_PARTNER_01" return m @@ -137,7 +137,7 @@ def test_get_transaction_detail_sftp(): mock_msg = MagicMock() mock_msg.id = uuid.uuid4() mock_msg.trace_id = uuid.uuid4() - mock_msg.outbound_route_id = uuid.uuid4() + mock_msg.trading_partner_id = "TEST_PARTNER_01" mock_msg.created_at = None mock_repo = AsyncMock() @@ -178,7 +178,7 @@ def test_get_transaction_detail_fallback(): mock_msg = MagicMock() mock_msg.id = uuid.uuid4() mock_msg.trace_id = uuid.uuid4() - mock_msg.outbound_route_id = None + mock_msg.trading_partner_id = None mock_msg.created_at = None mock_json = MagicMock() @@ -246,7 +246,7 @@ def test_get_transaction_webhook_fallback(): mock_msg.status = "RECEIVED" mock_msg.edi_data = "raw edi payload" mock_msg.created_at = None - mock_msg.outbound_route_id = None + mock_msg.trading_partner_id = None mock_json = MagicMock() mock_json.id = uuid.uuid4() diff --git a/services/workers/compute/src/compute_worker/main.py b/services/workers/compute/src/compute_worker/main.py index c30213f7..954197db 100644 --- a/services/workers/compute/src/compute_worker/main.py +++ b/services/workers/compute/src/compute_worker/main.py @@ -1,7 +1,12 @@ +# ruff: noqa: E402 import asyncio import logging import sys +from dotenv import load_dotenv + +load_dotenv() + from transformer.application.use_cases import ProcessInboundEdiUseCase from transformer.domain.models import ParsedEdiPayload from transformer.infrastructure.adapters.bots_adapter import BotsEDIAdapter diff --git a/services/workers/compute/src/compute_worker/worker.py b/services/workers/compute/src/compute_worker/worker.py index 81999096..0d461a1b 100644 --- a/services/workers/compute/src/compute_worker/worker.py +++ b/services/workers/compute/src/compute_worker/worker.py @@ -31,7 +31,7 @@ async def start(self) -> None: self._running = True logger.info(f"Starting SQS worker polling against {self.queue_url}") - client_kwargs = {} + client_kwargs = {"region_name": "us-east-1"} if self.endpoint_url: client_kwargs["endpoint_url"] = self.endpoint_url diff --git a/services/workers/orchestrator/pyproject.toml b/services/workers/orchestrator/pyproject.toml index 48e21b14..02b64848 100644 --- a/services/workers/orchestrator/pyproject.toml +++ b/services/workers/orchestrator/pyproject.toml @@ -13,6 +13,7 @@ dependencies = [ "pipeline", "config", "domain", + "scheduler", ] [tool.uv.sources] @@ -20,6 +21,7 @@ database = { workspace = true } pipeline = { workspace = true } config = { workspace = true } domain = { workspace = true } +scheduler = { workspace = true } [build-system] requires = ["hatchling"] diff --git a/services/workers/orchestrator/src/worker/adapters/db_outbox.py b/services/workers/orchestrator/src/worker/adapters/db_outbox.py index 2449e993..52912d42 100644 --- a/services/workers/orchestrator/src/worker/adapters/db_outbox.py +++ b/services/workers/orchestrator/src/worker/adapters/db_outbox.py @@ -4,7 +4,7 @@ from contextlib import asynccontextmanager from database.connection import DatabaseRouter -from database.models.control_plane import Outbox as GlobalOutbox +from database.models.control_plane import ControlPlaneOutbox from domain.events import ProvisioningEventType from sqlalchemy import select @@ -25,10 +25,10 @@ async def process_next_event(self) -> AsyncIterator[OutboxEvent | None]: try: stmt = ( - select(GlobalOutbox) + select(ControlPlaneOutbox) .where( - GlobalOutbox.status == "PENDING", - GlobalOutbox.event_type.in_(list(ProvisioningEventType)), + ControlPlaneOutbox.status == "PENDING", + ControlPlaneOutbox.event_type.in_(list(ProvisioningEventType)), ) .limit(1) .with_for_update(skip_locked=True) diff --git a/services/workers/orchestrator/src/worker/data/main.py b/services/workers/orchestrator/src/worker/data/main.py index 137029d5..e5784a27 100644 --- a/services/workers/orchestrator/src/worker/data/main.py +++ b/services/workers/orchestrator/src/worker/data/main.py @@ -27,8 +27,11 @@ WebhookDeliveryStrategy, ) from pipeline.core.transformation import InboundTransformService, OutboundTransformService +from scheduler.adapters.repository import SqlAlchemyJobRepository +from scheduler.core.service import SchedulerWorkerService from sqlalchemy import select from worker.adapters.vault import WorkerVaultAdapter +from worker.jobs.outbox_sweeper import DataPlaneOutboxSweeperJobHandler load_dotenv() @@ -362,6 +365,21 @@ async def main() -> None: ) ) + # Start Scheduler worker + from sqlalchemy.ext.asyncio import create_async_engine + + engine = create_async_engine(settings.database.global_url) + scheduler_repo = SqlAlchemyJobRepository(engine) + scheduler_service = SchedulerWorkerService( + scheduler_repo, worker_id=f"orchestrator-{os.getpid()}" + ) + scheduler_service.register_handler( + "outbox_sweeper", DataPlaneOutboxSweeperJobHandler(db_router) + ) + + # Run the scheduler loop in the background + await scheduler_service.start(poll_interval_seconds=10.0) + await asyncio.gather(transform_task, deliver_task) diff --git a/services/workers/orchestrator/src/worker/jobs/outbox_sweeper.py b/services/workers/orchestrator/src/worker/jobs/outbox_sweeper.py new file mode 100644 index 00000000..b7b565d4 --- /dev/null +++ b/services/workers/orchestrator/src/worker/jobs/outbox_sweeper.py @@ -0,0 +1,174 @@ +import asyncio +import datetime +import json +import logging +import os +from typing import Any + +import aioboto3 # type: ignore[import-untyped] +from database.connection import DatabaseRouter +from database.models.control_plane import DatabaseShard +from database.models.data_plane import DataPlaneOutbox +from domain.events import MessageQueueName, PipelineEventType +from scheduler.domain.models import Job +from scheduler.ports.handler import JobHandlerPort +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +logger = logging.getLogger(__name__) + +# Maps each pipeline event type to its target SQS queue +_EVENT_QUEUE_MAP: dict[str, str] = { + PipelineEventType.TRANSFORM_EVENT: MessageQueueName.TRANSFORM_ORCHESTRATION_QUEUE, + PipelineEventType.COMPUTE_TRANSFORM_EVENT: MessageQueueName.TRANSFORM_COMPUTE_QUEUE, + PipelineEventType.TRANSFORM_COMPLETED: MessageQueueName.TRANSFORM_ORCHESTRATION_QUEUE, + PipelineEventType.DELIVER_EVENT: MessageQueueName.DELIVER_QUEUE, + PipelineEventType.DELIVERY_COMPLETED: MessageQueueName.TRANSFORM_ORCHESTRATION_QUEUE, +} + +# Maximum number of events to sweep per run to bound wall-clock time +_BATCH_SIZE = 100 +_CONCURRENCY_LIMIT = 5 + + +class DataPlaneOutboxSweeperJobHandler(JobHandlerPort): + def __init__(self, db_router: DatabaseRouter) -> None: + self.db_router = db_router + self._endpoint_url = os.environ.get("AWS_ENDPOINT_URL", "http://localhost:4566") + self._region = "us-east-1" + self._session = aioboto3.Session() + + async def execute(self, job: Job) -> datetime.datetime | None: + """ + Sweeps the data-plane (tenant shard) outbox for PENDING pipeline events + and forwards each one to the appropriate SQS queue using concurrent batching. + """ + logger.info(f"[DataPlaneOutboxSweeper] Running sweep for job {job.id}") + + total_processed = 0 + sem = asyncio.Semaphore(_CONCURRENCY_LIMIT) + + # We share one SQS client pool across all shards + async with self._session.client( + "sqs", endpoint_url=self._endpoint_url, region_name=self._region + ) as sqs: + queue_url_cache: dict[str, str] = {} + + async for global_session in self.db_router.get_global_session(): + res = await global_session.execute(select(DatabaseShard)) + shards = res.scalars().all() + + async def _bounded_sweep(shard_name: str, shard_dsn: str) -> int: + async with sem: + return await self._sweep_shard(shard_name, shard_dsn, sqs, queue_url_cache) + + results = await asyncio.gather( + *[_bounded_sweep(shard.name, shard.dsn) for shard in shards] + ) + total_processed += sum(results) + + logger.info( + f"[DataPlaneOutboxSweeper] Sweep complete. Total events forwarded: {total_processed}" + ) + + interval_seconds = job.payload.get("interval_seconds", 60) if job.payload else 60 + return datetime.datetime.now(datetime.UTC) + datetime.timedelta(seconds=interval_seconds) + + async def _sweep_shard( + self, shard_name: str, shard_dsn: str, sqs: Any, queue_url_cache: dict[str, str] + ) -> int: + """Sweep a single tenant shard outbox, dispatching via SQS batching.""" + processed = 0 + + engine = await self.db_router.get_engine(shard_name, shard_dsn) + async with AsyncSession(engine, expire_on_commit=False) as session: + stmt = ( + select(DataPlaneOutbox) + .where( + DataPlaneOutbox.status == "PENDING", + DataPlaneOutbox.event_type.in_(list(PipelineEventType)), + ) + .limit(_BATCH_SIZE) + .with_for_update(skip_locked=True) + ) + result = await session.execute(stmt) + events = result.scalars().all() + + if not events: + logger.debug(f"[DataPlaneOutboxSweeper] No pending events on shard={shard_name}") + return 0 + + # Group events by target queue to utilize SQS send_message_batch + batches_by_queue: dict[str, list[DataPlaneOutbox]] = {} + for event in events: + queue_name = _EVENT_QUEUE_MAP.get(event.event_type) + if not queue_name: + logger.warning( + f"[DataPlaneOutboxSweeper] Unknown event_type={event.event_type!r} " + f"for event id={event.id}. Marking FAILED." + ) + event.status = "FAILED" + continue + batches_by_queue.setdefault(queue_name, []).append(event) + + for queue_name, queue_events in batches_by_queue.items(): + if queue_name not in queue_url_cache: + try: + resp = await sqs.get_queue_url(QueueName=queue_name) + queue_url_cache[queue_name] = resp["QueueUrl"] + except Exception: + logger.exception( + f"[DataPlaneOutboxSweeper] Failed to get queue url for {queue_name}" + ) + continue + + queue_url = queue_url_cache[queue_name] + + # SQS allows max 10 messages per batch + for i in range(0, len(queue_events), 10): + batch = queue_events[i : i + 10] + entries = [] + for event in batch: + entries.append( + { + "Id": str(event.id), # SQS entry ID must be string + "MessageBody": json.dumps( + { + "idempotency_key": str(event.idempotency_key), + "event_type": event.event_type, + "payload": event.payload, + "tenant_id": event.tenant_id, + } + ), + } + ) + + try: + resp = await sqs.send_message_batch(QueueUrl=queue_url, Entries=entries) + # Process successful IDs + for success in resp.get("Successful", []): + event_id = success["Id"] + # Find the event object + for ev in batch: + if str(ev.id) == event_id: + ev.status = "PROCESSED" + processed += 1 + break + + # Log failures if any + for failed in resp.get("Failed", []): + logger.error( + f"[DataPlaneOutboxSweeper] Failed to forward event id={failed['Id']}: " + f"{failed['Message']}" + ) + except Exception: + logger.exception( + f"[DataPlaneOutboxSweeper] Failed to send batch to {queue_name}" + ) + + # Commit the session. Only events marked PROCESSED/FAILED above will be updated. + # The ones that failed to send via SQS will simply remain PENDING in the DB (no modification made) + # because we did not change their status. + await session.commit() + + return processed diff --git a/uv.lock b/uv.lock index 8a052e85..dd7ccc09 100644 --- a/uv.lock +++ b/uv.lock @@ -25,6 +25,7 @@ members = [ "orchestrator-worker", "patches", "pipeline", + "scheduler", "security", "transformer", ] @@ -2134,6 +2135,7 @@ dependencies = [ { name = "database" }, { name = "domain" }, { name = "pipeline" }, + { name = "scheduler" }, { name = "sqlalchemy" }, ] @@ -2145,6 +2147,7 @@ requires-dist = [ { name = "database", editable = "libs/database" }, { name = "domain", editable = "libs/domain" }, { name = "pipeline", editable = "libs/pipeline" }, + { name = "scheduler", editable = "libs/scheduler" }, { name = "sqlalchemy", specifier = ">=2.0.29" }, ] @@ -2914,6 +2917,23 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/48/f0/ae7ca09223a81a1d890b2557186ea015f6e0502e9b8cb8e1813f1d8cfa4e/s3transfer-0.14.0-py3-none-any.whl", hash = "sha256:ea3b790c7077558ed1f02a3072fb3cb992bbbd253392f4b6e9e8976941c7d456", size = 85712, upload-time = "2025-09-09T19:23:30.041Z" }, ] +[[package]] +name = "scheduler" +version = "0.1.0" +source = { editable = "libs/scheduler" } +dependencies = [ + { name = "database" }, + { name = "pydantic" }, + { name = "sqlalchemy" }, +] + +[package.metadata] +requires-dist = [ + { name = "database", editable = "libs/database" }, + { name = "pydantic", specifier = ">=2.0.0" }, + { name = "sqlalchemy", specifier = ">=2.0.0" }, +] + [[package]] name = "security" version = "0.1.0" From fc057f5a23a1f5f0bfbcc0f81c0bb6e3a1f7087e Mon Sep 17 00:00:00 2001 From: Pramod Date: Fri, 17 Jul 2026 18:45:22 +0530 Subject: [PATCH 3/5] light weight scheduler and outbox sweeper code review --- docker/localstack/init-aws.sh | 4 + .../src/features/partners/api/partnerHooks.ts | 7 +- .../partners/components/As2PartnerDetails.tsx | 16 +- .../components/CreatePartnerModal.tsx | 14 +- frontend/web/src/features/partners/types.ts | 1 + .../features/partners/utils/certificate.ts | 46 ++ .../src/features/platform/api/schedulerApi.ts | 15 + .../features/platform/api/schedulerHooks.ts | 13 + .../platform/components/CronBuilder.tsx | 135 ++++++ .../components/SchedulerDashboard.tsx | 266 +++++++---- libs/config/src/config/settings.py | 13 + .../42c7e50a7b1c_global_initial_schema.py | 7 + .../src/database/models/scheduled_job.py | 19 +- libs/scheduler/pyproject.toml | 1 + .../src/scheduler/adapters/repository.py | 55 ++- libs/scheduler/src/scheduler/core/service.py | 110 +++-- libs/scheduler/src/scheduler/domain/models.py | 36 +- libs/scheduler/src/scheduler/ports/handler.py | 5 +- .../src/scheduler/ports/publisher.py | 5 + .../src/scheduler/ports/repository.py | 12 +- libs/scheduler/src/scheduler/registry.py | 42 ++ services/api/src/api/adapters/http/dtos.py | 3 + .../api/adapters/transaction_repository.py | 22 +- .../api/src/api/routers/platform/__init__.py | 6 +- .../api/src/api/routers/platform/scheduler.py | 204 ++++++--- .../trading_partners/platform/as2_partners.py | 60 ++- .../src/api/services/api_receiver_service.py | 20 +- services/api/tests/test_scheduler.py | 130 ++++++ services/as2_server/scripts/seed.py | 46 ++ .../src/worker/adapters/sqs_poller.py | 87 ++++ .../src/worker/adapters/sqs_publisher.py | 110 +++++ .../src/worker/core/job_registry.py | 16 + .../orchestrator/src/worker/core/security.py | 52 +++ .../src/worker/core/tenant_resolver.py | 37 ++ .../orchestrator/src/worker/data/handlers.py | 197 +++++++++ .../orchestrator/src/worker/data/main.py | 412 ++++-------------- .../src/worker/data/scheduled_jobs_handler.py | 38 ++ .../src/worker/jobs/data_retention.py | 83 ++++ .../src/worker/jobs/outbox_sweeper.py | 123 ++---- .../src/worker/ports/message_publisher.py | 22 + .../orchestrator/tests/test_data_main.py | 101 ++++- .../orchestrator/tests/test_sqs_publisher.py | 93 ++++ uv.lock | 14 + 43 files changed, 2037 insertions(+), 661 deletions(-) create mode 100644 frontend/web/src/features/partners/utils/certificate.ts create mode 100644 frontend/web/src/features/platform/components/CronBuilder.tsx create mode 100644 libs/scheduler/src/scheduler/ports/publisher.py create mode 100644 libs/scheduler/src/scheduler/registry.py create mode 100644 services/api/tests/test_scheduler.py create mode 100644 services/workers/orchestrator/src/worker/adapters/sqs_poller.py create mode 100644 services/workers/orchestrator/src/worker/adapters/sqs_publisher.py create mode 100644 services/workers/orchestrator/src/worker/core/job_registry.py create mode 100644 services/workers/orchestrator/src/worker/core/security.py create mode 100644 services/workers/orchestrator/src/worker/core/tenant_resolver.py create mode 100644 services/workers/orchestrator/src/worker/data/handlers.py create mode 100644 services/workers/orchestrator/src/worker/data/scheduled_jobs_handler.py create mode 100644 services/workers/orchestrator/src/worker/jobs/data_retention.py create mode 100644 services/workers/orchestrator/src/worker/ports/message_publisher.py create mode 100644 services/workers/orchestrator/tests/test_sqs_publisher.py diff --git a/docker/localstack/init-aws.sh b/docker/localstack/init-aws.sh index 1f4d1c2e..0fb609be 100755 --- a/docker/localstack/init-aws.sh +++ b/docker/localstack/init-aws.sh @@ -29,4 +29,8 @@ 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\\\"}\"}" +awslocal sqs create-queue --queue-name edi-orchestrator-jobs-DLQ +JOBS_DLQ_ARN=$(awslocal sqs get-queue-attributes --queue-url http://localhost:4566/000000000000/edi-orchestrator-jobs-DLQ --attribute-names QueueArn --query 'Attributes.QueueArn' --output text) +awslocal sqs create-queue --queue-name edi-orchestrator-jobs --attributes "{\"RedrivePolicy\":\"{\\\"deadLetterTargetArn\\\":\\\"$JOBS_DLQ_ARN\\\",\\\"maxReceiveCount\\\":\\\"3\\\"}\"}" + echo "LocalStack Initialization Complete." diff --git a/frontend/web/src/features/partners/api/partnerHooks.ts b/frontend/web/src/features/partners/api/partnerHooks.ts index dbabaa0c..b0097a64 100644 --- a/frontend/web/src/features/partners/api/partnerHooks.ts +++ b/frontend/web/src/features/partners/api/partnerHooks.ts @@ -138,10 +138,9 @@ export function useDeletePlatformPartnerMutation() { export function useDeleteCertificateSecretMutation() { const repo = useRepository(); - return useToastMutation( - (vaultRef: string) => repo.deleteCertificateSecret(vaultRef), - 'Orphaned certificate deleted successfully.' - ); + return useMutation({ + mutationFn: (vaultRef: string) => repo.deleteCertificateSecret(vaultRef), + }); } // ───────────────────────────────────────────── diff --git a/frontend/web/src/features/partners/components/As2PartnerDetails.tsx b/frontend/web/src/features/partners/components/As2PartnerDetails.tsx index 2f8d37c8..ba03bacd 100644 --- a/frontend/web/src/features/partners/components/As2PartnerDetails.tsx +++ b/frontend/web/src/features/partners/components/As2PartnerDetails.tsx @@ -12,6 +12,7 @@ import { useToast } from '@/hooks/use-toast'; import { usePlatformSettings } from '@/features/platform/api/settingsHooks'; import { Combobox } from '@/components/ui/combobox'; import { CertificateInput } from './CertificateInput'; +import { extractCertificateMaterial } from '../utils/certificate'; export function As2PartnerDetails({ partner, onCancel }: { partner: AS2Partner, onCancel?: () => void }) { const { toast } = useToast(); @@ -72,18 +73,7 @@ export function As2PartnerDetails({ partner, onCancel }: { partner: AS2Partner, const handlePasteSubmit = () => { if (!pasteValue.trim()) return; - let publicCert = ''; - let privateKey = ''; - - const text = pasteValue; - if (text.includes('-----BEGIN CERTIFICATE-----')) { - const match = text.match(/-----BEGIN CERTIFICATE-----[^-]+-----END CERTIFICATE-----/g); - if (match && match.length > 0) publicCert += match.join('\n') + '\n'; - } - if (text.includes('-----BEGIN PRIVATE KEY-----') || text.includes('-----BEGIN RSA PRIVATE KEY-----')) { - const match = text.match(/-----BEGIN (?:RSA )?PRIVATE KEY-----[^-]+-----END (?:RSA )?PRIVATE KEY-----/g); - if (match && match.length > 0) privateKey += match.join('\n') + '\n'; - } + const { publicCert, privateKey } = extractCertificateMaterial(pasteValue); rotateCertificates.mutate({ id: partner.id, @@ -292,7 +282,7 @@ function CertificateRow({ notBefore: cert.validity.notBefore.toLocaleDateString(undefined, { year: 'numeric', month: 'short', day: 'numeric' }), notAfter: cert.validity.notAfter.toLocaleDateString(undefined, { year: 'numeric', month: 'short', day: 'numeric' }), }; - } catch (_e) { + } catch { return null; } }, [publicPem]); diff --git a/frontend/web/src/features/partners/components/CreatePartnerModal.tsx b/frontend/web/src/features/partners/components/CreatePartnerModal.tsx index 3132ebd7..3ce2974e 100644 --- a/frontend/web/src/features/partners/components/CreatePartnerModal.tsx +++ b/frontend/web/src/features/partners/components/CreatePartnerModal.tsx @@ -9,6 +9,7 @@ import { useToast } from '@/hooks/use-toast'; import { Combobox } from '@/components/ui/combobox'; import { Button } from '@/components/ui/button'; import { Loader2 } from 'lucide-react'; +import { extractCertificateMaterial } from '../utils/certificate'; export function CreatePartnerModal({ existingAs2Ids = [] }: { existingAs2Ids?: string[] }) { const [isOpen, setIsOpen] = useState(false); @@ -76,6 +77,13 @@ export function CreatePartnerModal({ existingAs2Ids = [] }: { existingAs2Ids?: s // Check if AS2 ID changed after generating cert let finalCertPem = certPem; let finalVaultRef = privateKeyVaultRef; + let extractedPrivateKey = ''; + + if (isLocal && !privateKeyVaultRef && certPem) { + const { publicCert, privateKey } = extractCertificateMaterial(certPem); + finalCertPem = publicCert; + extractedPrivateKey = privateKey; + } if (isLocal && privateKeyVaultRef && generatedForAs2Id && submittedAs2Id !== generatedForAs2Id) { // Invalidate existing if AS2 ID changed @@ -96,8 +104,12 @@ export function CreatePartnerModal({ existingAs2Ids = [] }: { existingAs2Ids?: s as2_id: submittedAs2Id, is_local: isLocal, url: url, - public_cert_pem: isLocal && finalVaultRef ? finalCertPem : isLocal ? undefined : finalCertPem, + // Always pass the cert PEM if the user provided one + public_cert_pem: finalCertPem || undefined, + // If user went through the generate flow, send vault ref + // If user uploaded their own key, send the raw PEM so backend stores it private_key_vault_ref: finalVaultRef || undefined, + private_key_pem: (isLocal && !finalVaultRef && extractedPrivateKey) ? extractedPrivateKey : undefined, }, { onSuccess: () => { diff --git a/frontend/web/src/features/partners/types.ts b/frontend/web/src/features/partners/types.ts index c3a42497..69a6f75b 100644 --- a/frontend/web/src/features/partners/types.ts +++ b/frontend/web/src/features/partners/types.ts @@ -69,6 +69,7 @@ export interface CreatePartnerPayload { url?: string; public_cert_pem?: string; private_key_vault_ref?: string; + private_key_pem?: string; } export interface UpdatePartnerPayload { diff --git a/frontend/web/src/features/partners/utils/certificate.ts b/frontend/web/src/features/partners/utils/certificate.ts new file mode 100644 index 00000000..6eaa20b8 --- /dev/null +++ b/frontend/web/src/features/partners/utils/certificate.ts @@ -0,0 +1,46 @@ +/** + * Utility functions for parsing and extracting certificate material. + */ + +export interface ParsedCertificateMaterial { + publicCert: string; + privateKey: string; +} + +/** + * Extracts the public certificate(s) and private key(s) from a combined PEM string. + * @param pemText The raw text containing one or more PEM blocks. + * @returns An object containing the extracted public cert and private key (trimmed). + */ +export function extractCertificateMaterial(pemText: string): ParsedCertificateMaterial { + let publicCert = ''; + let privateKey = ''; + + if (!pemText) { + return { publicCert, privateKey }; + } + + // Extract Public Certificates + if (pemText.includes('-----BEGIN CERTIFICATE-----')) { + const match = pemText.match(/-----BEGIN CERTIFICATE-----[^-]+-----END CERTIFICATE-----/g); + if (match && match.length > 0) { + publicCert = match.join('\n') + '\n'; + } + } + + // Extract Private Keys (supports both standard and RSA specific headers) + if ( + pemText.includes('-----BEGIN PRIVATE KEY-----') || + pemText.includes('-----BEGIN RSA PRIVATE KEY-----') + ) { + const match = pemText.match(/-----BEGIN (?:RSA )?PRIVATE KEY-----[^-]+-----END (?:RSA )?PRIVATE KEY-----/g); + if (match && match.length > 0) { + privateKey = match.join('\n') + '\n'; + } + } + + return { + publicCert: publicCert.trim(), + privateKey: privateKey.trim(), + }; +} diff --git a/frontend/web/src/features/platform/api/schedulerApi.ts b/frontend/web/src/features/platform/api/schedulerApi.ts index d23ce7da..d584c363 100644 --- a/frontend/web/src/features/platform/api/schedulerApi.ts +++ b/frontend/web/src/features/platform/api/schedulerApi.ts @@ -2,10 +2,17 @@ export interface JobResponse { id: string; name: string; status: string; + target_queue: string | null; + app_namespace: string | null; + cron_expression: string | null; + timezone: string | null; next_run_at: string | null; locked_at: string | null; locked_by: string | null; retry_count: number; + interval_seconds: number | null; + min_interval_seconds: number | null; + max_interval_seconds: number | null; error_message: string | null; created_at: string; updated_at: string; @@ -20,6 +27,7 @@ export interface ISchedulerRepository { getJobs(): Promise; getConfig(): Promise; updateConfig(key: string, value: any): Promise; + updateJob(name: string, data: { interval_seconds?: number | null; cron_expression?: string | null; timezone?: string | null; status?: string }): Promise; } class HttpSchedulerRepository implements ISchedulerRepository { @@ -61,6 +69,13 @@ class HttpSchedulerRepository implements ISchedulerRepository { body: JSON.stringify({ value }), }); } + + updateJob(name: string, data: { interval_seconds?: number | null; cron_expression?: string | null; timezone?: string | null; status?: string }): Promise { + return this.request(`/api/v1/platform/scheduler/jobs/${name}`, { + method: 'PUT', + body: JSON.stringify(data), + }); + } } export function createSchedulerRepository(token: string): ISchedulerRepository { diff --git a/frontend/web/src/features/platform/api/schedulerHooks.ts b/frontend/web/src/features/platform/api/schedulerHooks.ts index 7a76b764..458ad2f3 100644 --- a/frontend/web/src/features/platform/api/schedulerHooks.ts +++ b/frontend/web/src/features/platform/api/schedulerHooks.ts @@ -37,3 +37,16 @@ export function useUpdateConfigMutation() { }, }); } + +export function useUpdateJobMutation() { + const repo = useRepo(); + const queryClient = useQueryClient(); + + return useMutation({ + mutationFn: ({ name, data }: { name: string; data: { interval_seconds?: number | null; cron_expression?: string | null; timezone?: string | null; status?: string } }) => + repo.updateJob(name, data), + onSuccess: () => { + queryClient.invalidateQueries({ queryKey: ['scheduler', 'jobs'] }); + }, + }); +} diff --git a/frontend/web/src/features/platform/components/CronBuilder.tsx b/frontend/web/src/features/platform/components/CronBuilder.tsx new file mode 100644 index 00000000..f5c0e769 --- /dev/null +++ b/frontend/web/src/features/platform/components/CronBuilder.tsx @@ -0,0 +1,135 @@ +import { useState, useEffect } from 'react'; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from '@/components/ui/select'; +import { Label } from '@/components/ui/label'; + +interface CronBuilderProps { + value: string; + onChange: (cronString: string) => void; +} + +export function CronBuilder({ value, onChange }: CronBuilderProps) { + // Parse initial value (defaults to '0 * * * *' if invalid) + const parts = value.split(' '); + // If 5 parts, prepend '0' for seconds. If 6 parts, use as is. Otherwise default 6 parts. + const initial = parts.length === 6 ? parts : parts.length === 5 ? ['0', ...parts] : ['0', '0', '*', '*', '*', '*']; + + const [second, setSecond] = useState(initial[0]); + const [minute, setMinute] = useState(initial[1]); + const [hour, setHour] = useState(initial[2]); + const [dayOfMonth, setDayOfMonth] = useState(initial[3]); + const [month, setMonth] = useState(initial[4]); + const [dayOfWeek, setDayOfWeek] = useState(initial[5]); + + useEffect(() => { + onChange(`${second} ${minute} ${hour} ${dayOfMonth} ${month} ${dayOfWeek}`); + }, [second, minute, hour, dayOfMonth, month, dayOfWeek, onChange]); + + const generateOptions = (start: number, end: number, labelPrefix = '') => { + const opts = [{ value: '*', label: 'Every' }]; + for (let i = start; i <= end; i++) { + opts.push({ value: String(i), label: `${labelPrefix}${i}` }); + } + return opts; + }; + + const seconds = generateOptions(0, 59); + const minutes = generateOptions(0, 59); + const hours = generateOptions(0, 23); + const daysOfMonth = generateOptions(1, 31); + const months = [ + { value: '*', label: 'Every' }, + { value: '1', label: 'January' }, + { value: '2', label: 'February' }, + { value: '3', label: 'March' }, + { value: '4', label: 'April' }, + { value: '5', label: 'May' }, + { value: '6', label: 'June' }, + { value: '7', label: 'July' }, + { value: '8', label: 'August' }, + { value: '9', label: 'September' }, + { value: '10', label: 'October' }, + { value: '11', label: 'November' }, + { value: '12', label: 'December' }, + ]; + const daysOfWeek = [ + { value: '*', label: 'Every' }, + { value: '0', label: 'Sunday' }, + { value: '1', label: 'Monday' }, + { value: '2', label: 'Tuesday' }, + { value: '3', label: 'Wednesday' }, + { value: '4', label: 'Thursday' }, + { value: '5', label: 'Friday' }, + { value: '6', label: 'Saturday' }, + ]; + + return ( +
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ ); +} diff --git a/frontend/web/src/features/platform/components/SchedulerDashboard.tsx b/frontend/web/src/features/platform/components/SchedulerDashboard.tsx index 856e7c22..e710270d 100644 --- a/frontend/web/src/features/platform/components/SchedulerDashboard.tsx +++ b/frontend/web/src/features/platform/components/SchedulerDashboard.tsx @@ -1,52 +1,100 @@ -import { useState, useEffect } from 'react'; -import { useJobsQuery, useConfigQuery, useUpdateConfigMutation } from '../api/schedulerHooks'; +import { useState } from 'react'; +import { useJobsQuery, useUpdateJobMutation } from '../api/schedulerHooks'; import { Card, CardContent, CardHeader, CardTitle } from '@/components/ui/card'; import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from '@/components/ui/table'; -import { Power } from 'lucide-react'; -import { Label } from '@/components/ui/label'; -import { Input } from '@/components/ui/input'; import { Button } from '@/components/ui/button'; +import { Input } from '@/components/ui/input'; +import { Label } from '@/components/ui/label'; +import { Pause, Play, Clock } from 'lucide-react'; +import { + Dialog, + DialogContent, + DialogDescription, + DialogFooter, + DialogHeader, + DialogTitle, +} from "@/components/ui/dialog"; +import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs'; +import { CronBuilder } from './CronBuilder'; export const SchedulerDashboard = () => { - const { data: jobs = [], isLoading: jobsLoading } = useJobsQuery(); - const { data: configs = [], isLoading: configLoading } = useConfigQuery(); - const { mutateAsync: updateConfig } = useUpdateConfigMutation(); - const getSweeperEnabled = () => { - const cfg = configs.find((c: any) => c.key === 'outbox_sweeper_enabled'); - return cfg ? !!cfg.value : false; - }; + const { data: jobs = [], isLoading: jobsLoading, refetch: refetchJobs } = useJobsQuery(); + const { mutateAsync: updateJob } = useUpdateJobMutation(); + + const [editingJob, setEditingJob] = useState(null); + const [scheduleType, setScheduleType] = useState<'interval' | 'cron'>('interval'); + const [intervalValue, setIntervalValue] = useState('1'); + const [intervalUnit, setIntervalUnit] = useState('minutes'); + const [newCron, setNewCron] = useState('0 * * * *'); - const getSweeperInterval = () => { - const cfg = configs.find((c: any) => c.key === 'outbox_sweeper_interval_seconds'); - return cfg ? cfg.value : 60; + const handleTogglePause = async (job: any) => { + const newStatus = job.status === 'PAUSED' ? 'PENDING' : 'PAUSED'; + try { + await updateJob({ name: job.name, data: { status: newStatus } }); + } catch (e) { + console.error('Failed to toggle job status', e); + } }; - const [localEnabled, setLocalEnabled] = useState(false); - const [localInterval, setLocalInterval] = useState(''); + const [intervalError, setIntervalError] = useState(null); + + const getIntervalParts = (seconds: number) => { + if (seconds % (30 * 24 * 3600) === 0) return { val: seconds / (30 * 24 * 3600), unit: 'months' }; + if (seconds % (24 * 3600) === 0) return { val: seconds / (24 * 3600), unit: 'days' }; + if (seconds % 3600 === 0) return { val: seconds / 3600, unit: 'hours' }; + if (seconds % 60 === 0) return { val: seconds / 60, unit: 'minutes' }; + return { val: seconds, unit: 'seconds' }; + }; - useEffect(() => { - if (configs.length > 0) { - setLocalEnabled(getSweeperEnabled()); - setLocalInterval(String(getSweeperInterval())); + const handleEditIntervalClick = (job: any) => { + setEditingJob(job); + if (job.cron_expression) { + setScheduleType('cron'); + setNewCron(job.cron_expression); + const parts = getIntervalParts(job.interval_seconds || 60); + setIntervalValue(String(parts.val)); + setIntervalUnit(parts.unit); + } else { + setScheduleType('interval'); + const parts = getIntervalParts(job.interval_seconds || 60); + setIntervalValue(String(parts.val)); + setIntervalUnit(parts.unit); + setNewCron('0 * * * *'); } - }, [configs, getSweeperEnabled, getSweeperInterval]); + setIntervalError(null); + }; + + const handleSaveSchedule = async () => { + if (!editingJob) return; - const handleSave = async () => { try { - if (localInterval) { - await updateConfig({ key: 'outbox_sweeper_interval_seconds', value: parseInt(localInterval, 10) }); + if (scheduleType === 'interval') { + const val = parseInt(intervalValue, 10); + if (isNaN(val) || val <= 0) { + setIntervalError('Interval must be a positive integer.'); + return; + } + + let multiplier = 1; + if (intervalUnit === 'minutes') multiplier = 60; + else if (intervalUnit === 'hours') multiplier = 3600; + else if (intervalUnit === 'days') multiplier = 86400; + else if (intervalUnit === 'months') multiplier = 2592000; + + const totalSeconds = val * multiplier; + + await updateJob({ name: editingJob.name, data: { interval_seconds: totalSeconds, cron_expression: null, timezone: null } }); + } else { + await updateJob({ name: editingJob.name, data: { cron_expression: newCron, timezone: 'UTC', interval_seconds: null } }); } - await updateConfig({ key: 'outbox_sweeper_enabled', value: localEnabled }); - } catch (e) { - console.error('Failed to update config', e); + setEditingJob(null); + setIntervalError(null); + } catch (e: any) { + setIntervalError(e.message || 'Failed to update job schedule'); } }; - const isDirty = localEnabled !== getSweeperEnabled() || localInterval !== String(getSweeperInterval()); - const isValid = localInterval.trim() !== ''; - const canSave = isDirty && isValid; - - if (jobsLoading || configLoading) return
Loading Scheduler Dashboard...
; + if (jobsLoading) return
Loading Scheduler Dashboard...
; return (
@@ -56,61 +104,12 @@ export const SchedulerDashboard = () => {
Scheduled Jobs - {getSweeperEnabled() ? ( -

- {(() => { - const secs = getSweeperInterval(); - return secs < 60 - ? `Outbox Sweeper is scheduled to run every ${secs}s.` - : `Outbox Sweeper is scheduled to run every ${Math.round(secs / 60)} min${Math.round(secs / 60) !== 1 ? 's' : ''}.`; - })()} -

- ) : ( -

- Outbox Sweeper is not currently scheduled. -

- )} +

+ Overview of all background jobs running in the platform. +

-
- - -
- -
- - setLocalInterval(e.target.value)} - disabled={!localEnabled} - className="w-20 h-8" - /> - -
- -
@@ -119,38 +118,50 @@ export const SchedulerDashboard = () => { - Job ID Name Status - Error Message + Schedule Next Run At - Locked By + Actions {jobs.map((job: any) => ( - {job.id} - {job.name} + {job.name} {job.status} - - {job.error_message || '-'} + + {job.cron_expression + ? {job.cron_expression} + : job.interval_seconds ? `${job.interval_seconds}s` : '-'} {job.next_run_at ? new Date(job.next_run_at).toLocaleString() : 'Immediate'} - {job.locked_by || '-'} + +
+ + +
+
))} {jobs.length === 0 && ( - + No background jobs found. @@ -159,6 +170,69 @@ export const SchedulerDashboard = () => {
+ + !open && setEditingJob(null)}> + + + Edit Job Schedule + + Set how often the {editingJob?.name} job runs. + {editingJob?.min_interval_seconds && editingJob?.max_interval_seconds && ( + + Allowed range (if using interval): {editingJob.min_interval_seconds}s - {editingJob.max_interval_seconds}s + + )} + + + + setScheduleType(v as any)}> + + Interval + Cron Expression + + +
+ + { setIntervalValue(e.target.value); setIntervalError(null); }} + className={`w-24 ${intervalError ? 'border-red-500' : ''}`} + /> + +
+
+ + { setNewCron(val); setIntervalError(null); }} /> +
+ Generated Cron: + {newCron} +
+
+
+ + {intervalError && ( +

{intervalError}

+ )} + + + + +
+
); }; diff --git a/libs/config/src/config/settings.py b/libs/config/src/config/settings.py index 7810d981..49f6b6b5 100644 --- a/libs/config/src/config/settings.py +++ b/libs/config/src/config/settings.py @@ -37,6 +37,18 @@ class S3Settings(BaseSettings): secret_access_key: str | None = Field(default=None) +class AwsSettings(BaseSettings): + model_config = SettingsConfigDict(env_prefix="AWS_", env_file=".env", extra="ignore") + + endpoint_url: str | None = Field(default=None) + region: str | None = Field(default=None) + default_region: str = Field(default="us-east-1", validation_alias="AWS_DEFAULT_REGION") + + @property + def resolved_region(self) -> str: + return self.region or self.default_region + + class OtelSettings(BaseSettings): model_config = SettingsConfigDict(env_prefix="OTEL_", env_file=".env", extra="ignore") @@ -106,6 +118,7 @@ class AppSettings(BaseSettings): database: DatabaseSettings = Field(default_factory=DatabaseSettings) s3: S3Settings = Field(default_factory=S3Settings) + aws: AwsSettings = Field(default_factory=AwsSettings) otel: OtelSettings = Field(default_factory=OtelSettings) identity: IdentitySettings = Field(default_factory=IdentitySettings) server: ServerSettings = Field(default_factory=ServerSettings) diff --git a/libs/database/src/database/migrations/global/versions/42c7e50a7b1c_global_initial_schema.py b/libs/database/src/database/migrations/global/versions/42c7e50a7b1c_global_initial_schema.py index 3095a393..3c165642 100644 --- a/libs/database/src/database/migrations/global/versions/42c7e50a7b1c_global_initial_schema.py +++ b/libs/database/src/database/migrations/global/versions/42c7e50a7b1c_global_initial_schema.py @@ -356,6 +356,13 @@ def upgrade() -> None: sa.Column("next_run_at", sa.DateTime(timezone=True), nullable=True), sa.Column("locked_at", sa.DateTime(timezone=True), nullable=True), sa.Column("locked_by", sa.String(), nullable=True), + sa.Column("target_queue", sa.String(), nullable=True), + sa.Column("app_namespace", sa.String(), nullable=True), + sa.Column("cron_expression", sa.String(), nullable=True), + sa.Column("timezone", sa.String(), nullable=True), + sa.Column("interval_seconds", sa.Integer(), nullable=True), + sa.Column("min_interval_seconds", sa.Integer(), nullable=True), + sa.Column("max_interval_seconds", sa.Integer(), nullable=True), sa.Column("retry_count", sa.Integer(), nullable=False), sa.Column("max_retries", sa.Integer(), nullable=False), sa.Column("error_message", sa.String(), nullable=True), diff --git a/libs/database/src/database/models/scheduled_job.py b/libs/database/src/database/models/scheduled_job.py index 48f1c74e..a0ce0957 100644 --- a/libs/database/src/database/models/scheduled_job.py +++ b/libs/database/src/database/models/scheduled_job.py @@ -1,5 +1,5 @@ import uuid -from datetime import datetime +from datetime import UTC, datetime from typing import Any from sqlalchemy import JSON, DateTime, Integer, String @@ -16,17 +16,30 @@ class ScheduledJob(GlobalBase): payload: Mapped[dict[str, Any]] = mapped_column(JSON, default=dict) status: Mapped[str] = mapped_column(String, index=True, default="PENDING") + target_queue: Mapped[str | None] = mapped_column(String, nullable=True) + app_namespace: Mapped[str | None] = mapped_column(String, nullable=True) + cron_expression: Mapped[str | None] = mapped_column(String, nullable=True) + timezone: Mapped[str | None] = mapped_column(String, nullable=True) + next_run_at: Mapped[datetime | None] = mapped_column( DateTime(timezone=True), nullable=True, index=True ) locked_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True) locked_by: Mapped[str | None] = mapped_column(String, nullable=True) + interval_seconds: Mapped[int | None] = mapped_column(Integer, nullable=True) + min_interval_seconds: Mapped[int | None] = mapped_column(Integer, nullable=True) + max_interval_seconds: Mapped[int | None] = mapped_column(Integer, nullable=True) + retry_count: Mapped[int] = mapped_column(Integer, default=0) max_retries: Mapped[int] = mapped_column(Integer, default=3) error_message: Mapped[str | None] = mapped_column(String, nullable=True) - created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=datetime.utcnow) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), default=lambda: datetime.now(UTC) + ) updated_at: Mapped[datetime] = mapped_column( - DateTime(timezone=True), default=datetime.utcnow, onupdate=datetime.utcnow + DateTime(timezone=True), + default=lambda: datetime.now(UTC), + onupdate=lambda: datetime.now(UTC), ) diff --git a/libs/scheduler/pyproject.toml b/libs/scheduler/pyproject.toml index bf0efc34..eca9ef56 100644 --- a/libs/scheduler/pyproject.toml +++ b/libs/scheduler/pyproject.toml @@ -9,6 +9,7 @@ dependencies = [ "SQLAlchemy>=2.0.0", # Local dependencies "database", + "croniter>=6.2.4", ] [tool.uv.sources] diff --git a/libs/scheduler/src/scheduler/adapters/repository.py b/libs/scheduler/src/scheduler/adapters/repository.py index e21297a0..e7f151fa 100644 --- a/libs/scheduler/src/scheduler/adapters/repository.py +++ b/libs/scheduler/src/scheduler/adapters/repository.py @@ -23,6 +23,11 @@ def _to_domain(self, record: ScheduledJob) -> Job: payload=record.payload, status=JobStatus(record.status), next_run_at=record.next_run_at, + target_queue=record.target_queue, + app_namespace=record.app_namespace, + cron_expression=record.cron_expression, + timezone=record.timezone, + interval_seconds=record.interval_seconds, retry_count=record.retry_count, max_retries=record.max_retries, locked_at=record.locked_at, @@ -31,9 +36,9 @@ def _to_domain(self, record: ScheduledJob) -> Job: updated_at=record.updated_at, ) - async def claim_next_job(self, worker_id: str) -> Job | None: + async def claim_next_jobs(self, worker_id: str, limit: int) -> list[Job]: """ - Uses SKIP LOCKED to safely claim the next PENDING or ready job. + Uses SKIP LOCKED to safely claim the next PENDING or ready jobs. """ now = datetime.datetime.now(datetime.UTC) @@ -45,24 +50,28 @@ async def claim_next_job(self, worker_id: str) -> Job | None: (ScheduledJob.status == JobStatus.PENDING.value) & (ScheduledJob.next_run_at.is_(None) | (ScheduledJob.next_run_at <= now)) ) - .order_by(ScheduledJob.created_at.asc()) - .limit(1) + .order_by( + ScheduledJob.next_run_at.asc().nulls_first(), ScheduledJob.created_at.asc() + ) + .limit(limit) .with_for_update(skip_locked=True) ) result = await session.execute(stmt) - record = result.scalar_one_or_none() + records = result.scalars().all() - if not record: - return None + if not records: + return [] - # Claim it - record.status = JobStatus.RUNNING.value - record.locked_at = now - record.locked_by = worker_id + claimed = [] + for record in records: + record.status = JobStatus.RUNNING.value + record.locked_at = now + record.locked_by = worker_id + claimed.append(self._to_domain(record)) await session.flush() - return self._to_domain(record) + return claimed async def mark_completed(self, job_id: uuid.UUID) -> None: async with self.session_factory() as session, session.begin(): @@ -92,7 +101,11 @@ async def mark_failed(self, job_id: uuid.UUID, error: str) -> None: await session.execute(stmt) async def schedule_job( - self, name: str, payload: dict[str, Any], next_run_at: datetime.datetime | None = None + self, + name: str, + payload: dict[str, Any], + next_run_at: datetime.datetime | None = None, + interval_seconds: int | None = None, ) -> Job: async with self.session_factory() as session, session.begin(): record = ScheduledJob( @@ -100,6 +113,7 @@ async def schedule_job( payload=payload, status=JobStatus.PENDING.value, next_run_at=next_run_at, + interval_seconds=interval_seconds, ) session.add(record) await session.flush() @@ -119,3 +133,18 @@ async def reschedule(self, job_id: uuid.UUID, next_run_at: datetime.datetime) -> ) ) await session.execute(stmt) + + async def schedule_retry(self, job_id: uuid.UUID, next_run_at: datetime.datetime) -> None: + async with self.session_factory() as session, session.begin(): + stmt = ( + update(ScheduledJob) + .where(ScheduledJob.id == job_id) + .values( + status=JobStatus.PENDING.value, + next_run_at=next_run_at, + locked_at=None, + locked_by=None, + retry_count=ScheduledJob.retry_count + 1, + ) + ) + await session.execute(stmt) diff --git a/libs/scheduler/src/scheduler/core/service.py b/libs/scheduler/src/scheduler/core/service.py index 4dd69ac1..d3f533ab 100644 --- a/libs/scheduler/src/scheduler/core/service.py +++ b/libs/scheduler/src/scheduler/core/service.py @@ -1,27 +1,35 @@ import asyncio import contextlib import logging +from typing import Any -from scheduler.ports.handler import JobHandlerPort +from scheduler.ports.publisher import MessagePublisherPort from scheduler.ports.repository import JobRepositoryPort logger = logging.getLogger(__name__) class SchedulerWorkerService: - def __init__(self, repository: JobRepositoryPort, worker_id: str): + def __init__( + self, + repository: JobRepositoryPort, + publisher: MessagePublisherPort, + worker_id: str, + max_concurrent_jobs: int = 10, + ): self.repository = repository + self.publisher = publisher self.worker_id = worker_id - self.handlers: dict[str, JobHandlerPort] = {} + self.max_concurrent_jobs = max_concurrent_jobs self._is_running = False self._task: asyncio.Task[None] | None = None - - def register_handler(self, job_name: str, handler: JobHandlerPort) -> None: - self.handlers[job_name] = handler + self._active_jobs: set[asyncio.Task[None]] = set() async def start(self, poll_interval_seconds: float = 5.0) -> None: self._is_running = True - logger.info(f"Starting scheduler worker {self.worker_id}") + logger.info( + f"Starting scheduler worker {self.worker_id} with concurrency {self.max_concurrent_jobs}" + ) self._task = asyncio.create_task(self._poll_loop(poll_interval_seconds)) async def stop(self) -> None: @@ -30,36 +38,78 @@ async def stop(self) -> None: self._task.cancel() with contextlib.suppress(asyncio.CancelledError): await self._task + + if self._active_jobs: + logger.info(f"Waiting for {len(self._active_jobs)} active jobs to complete...") + await asyncio.gather(*self._active_jobs, return_exceptions=True) + logger.info(f"Stopped scheduler worker {self.worker_id}") + async def _execute_job(self, job: Any) -> None: + try: + if not job.target_queue: + error_msg = f"No target_queue defined for job {job.name}" + logger.error(error_msg) + await self.repository.mark_failed(job.id, error=error_msg) + return + + logger.info(f"Dispatching job {job.name} ({job.id}) to queue {job.target_queue}") + payload = { + "job_id": str(job.id), + "job_name": job.name, + "payload": job.payload, + } + await self.publisher.publish(job.target_queue, payload) + + # Successfully dispatched. Reschedule + import datetime + + now = datetime.datetime.now(datetime.UTC) + next_run_at = job.calculate_next_run_at(now) + + if next_run_at: + await self.repository.reschedule(job.id, next_run_at) + logger.info(f"Successfully rescheduled job {job.name} ({job.id}) for {next_run_at}") + else: + await self.repository.mark_completed(job.id) + logger.info(f"Successfully completed job {job.name} ({job.id})") + + except Exception as e: + logger.exception(f"Job {job.name} ({job.id}) dispatch failed: {e}") + if job.retry_count < job.max_retries: + import datetime + + backoff_seconds = 60 * (2**job.retry_count) + next_run_at = datetime.datetime.now(datetime.UTC) + datetime.timedelta( + seconds=backoff_seconds + ) + await self.repository.schedule_retry(job.id, next_run_at) + logger.info(f"Scheduled retry for job {job.name} ({job.id}) at {next_run_at}") + else: + await self.repository.mark_failed(job.id, error=str(e)) + async def _poll_loop(self, poll_interval_seconds: float) -> None: while self._is_running: try: - job = await self.repository.claim_next_job(worker_id=self.worker_id) - if job: - handler = self.handlers.get(job.name) - if not handler: - error_msg = f"No handler registered for job {job.name}" - logger.error(error_msg) - await self.repository.mark_failed(job.id, error=error_msg) - continue - - try: - logger.info(f"Executing job {job.name} ({job.id})") - next_run_at = await handler.execute(job) - if next_run_at: - await self.repository.reschedule(job.id, next_run_at) - logger.info( - f"Successfully rescheduled job {job.name} ({job.id}) for {next_run_at}" - ) - else: - await self.repository.mark_completed(job.id) - logger.info(f"Successfully completed job {job.name} ({job.id})") - except Exception as e: - logger.exception(f"Job {job.name} ({job.id}) failed: {e}") - await self.repository.mark_failed(job.id, error=str(e)) + # Remove completed tasks from active set + self._active_jobs = {task for task in self._active_jobs if not task.done()} + + available_slots = self.max_concurrent_jobs - len(self._active_jobs) + + if available_slots > 0: + jobs = await self.repository.claim_next_jobs( + worker_id=self.worker_id, limit=available_slots + ) + + for job in jobs: + task = asyncio.create_task(self._execute_job(job)) + self._active_jobs.add(task) + + if not jobs: + await asyncio.sleep(poll_interval_seconds) else: await asyncio.sleep(poll_interval_seconds) + except asyncio.CancelledError: break except Exception as e: diff --git a/libs/scheduler/src/scheduler/domain/models.py b/libs/scheduler/src/scheduler/domain/models.py index b256ea51..8244ff74 100644 --- a/libs/scheduler/src/scheduler/domain/models.py +++ b/libs/scheduler/src/scheduler/domain/models.py @@ -2,7 +2,7 @@ from dataclasses import dataclass, field from datetime import datetime from enum import StrEnum -from typing import Any +from typing import Any, cast class JobStatus(StrEnum): @@ -10,6 +10,12 @@ class JobStatus(StrEnum): RUNNING = "RUNNING" COMPLETED = "COMPLETED" FAILED = "FAILED" + PAUSED = "PAUSED" + + +class JobName(StrEnum): + OUTBOX_SWEEPER = "outbox_sweeper" + DATA_RETENTION_CLEANUP = "data_retention_cleanup" @dataclass @@ -17,11 +23,39 @@ class Job: name: str payload: dict[str, Any] status: JobStatus = JobStatus.PENDING + target_queue: str | None = None + app_namespace: str | None = None next_run_at: datetime | None = None + interval_seconds: int | None = None + cron_expression: str | None = None + timezone: str | None = None retry_count: int = 0 max_retries: int = 3 locked_at: datetime | None = None locked_by: str | None = None + + def calculate_next_run_at(self, now: datetime) -> datetime | None: + """ + Calculate the next run time based on cron_expression or interval_seconds. + If neither is provided, returns None. + """ + import datetime as dt + + if self.cron_expression: + import zoneinfo + + from croniter import croniter # type: ignore + + tz = zoneinfo.ZoneInfo(self.timezone) if self.timezone else dt.UTC + now_tz = now.astimezone(tz) + itr = croniter(self.cron_expression, now_tz, second_at_beginning=True) + return cast(dt.datetime, itr.get_next(dt.datetime)) + + if self.interval_seconds: + return now + dt.timedelta(seconds=self.interval_seconds) + + return None + id: uuid.UUID = field(default_factory=uuid.uuid4) created_at: datetime | None = None updated_at: datetime | None = None diff --git a/libs/scheduler/src/scheduler/ports/handler.py b/libs/scheduler/src/scheduler/ports/handler.py index 8642653c..1da86ef6 100644 --- a/libs/scheduler/src/scheduler/ports/handler.py +++ b/libs/scheduler/src/scheduler/ports/handler.py @@ -1,11 +1,10 @@ import abc -import datetime from scheduler.domain.models import Job class JobHandlerPort(abc.ABC): @abc.abstractmethod - async def execute(self, job: Job) -> datetime.datetime | None: - """Execute the job. Return a datetime to reschedule it, or None to complete it.""" + async def execute(self, job: Job) -> None: + """Execute the job. Handlers should raise exceptions on failure.""" pass diff --git a/libs/scheduler/src/scheduler/ports/publisher.py b/libs/scheduler/src/scheduler/ports/publisher.py new file mode 100644 index 00000000..9356e160 --- /dev/null +++ b/libs/scheduler/src/scheduler/ports/publisher.py @@ -0,0 +1,5 @@ +import typing + + +class MessagePublisherPort(typing.Protocol): + async def publish(self, queue_name: str, payload: dict[str, typing.Any]) -> None: ... diff --git a/libs/scheduler/src/scheduler/ports/repository.py b/libs/scheduler/src/scheduler/ports/repository.py index a4647986..16c88a11 100644 --- a/libs/scheduler/src/scheduler/ports/repository.py +++ b/libs/scheduler/src/scheduler/ports/repository.py @@ -8,7 +8,7 @@ class JobRepositoryPort(abc.ABC): @abc.abstractmethod - async def claim_next_job(self, worker_id: str) -> Job | None: + async def claim_next_jobs(self, worker_id: str, limit: int) -> list[Job]: pass @abc.abstractmethod @@ -21,10 +21,18 @@ async def mark_failed(self, job_id: uuid.UUID, error: str) -> None: @abc.abstractmethod async def schedule_job( - self, name: str, payload: dict[str, Any], next_run_at: datetime | None = None + self, + name: str, + payload: dict[str, Any], + next_run_at: datetime | None = None, + interval_seconds: int | None = None, ) -> Job: pass @abc.abstractmethod async def reschedule(self, job_id: uuid.UUID, next_run_at: datetime) -> None: pass + + @abc.abstractmethod + async def schedule_retry(self, job_id: uuid.UUID, next_run_at: datetime) -> None: + pass diff --git a/libs/scheduler/src/scheduler/registry.py b/libs/scheduler/src/scheduler/registry.py new file mode 100644 index 00000000..7ee76977 --- /dev/null +++ b/libs/scheduler/src/scheduler/registry.py @@ -0,0 +1,42 @@ +from dataclasses import dataclass + +from scheduler.domain.models import JobName + + +@dataclass(frozen=True) +class JobDefinition: + """ + Canonical definition of a platform system job. + Drives both database seeding and API-level validation of job configuration. + """ + + name: JobName + target_queue: str | None = None + app_namespace: str | None = None + default_interval_seconds: int | None = None + min_interval_seconds: int | None = None + max_interval_seconds: int | None = None + default_cron_expression: str | None = None + default_timezone: str | None = None + max_retries: int = 3 + + +# Central registry of all platform-managed background jobs. +# To register a new system job, append a JobDefinition here. +SYSTEM_JOB_REGISTRY: list[JobDefinition] = [ + JobDefinition( + name=JobName.OUTBOX_SWEEPER, + target_queue="edi-orchestrator-jobs", + app_namespace="EDI", + default_interval_seconds=60, + min_interval_seconds=10, + max_interval_seconds=300, + ), + JobDefinition( + name=JobName.DATA_RETENTION_CLEANUP, + target_queue="edi-orchestrator-jobs", + app_namespace="EDI", + default_cron_expression="0 2 * * *", # 2 AM daily + default_timezone="UTC", + ), +] diff --git a/services/api/src/api/adapters/http/dtos.py b/services/api/src/api/adapters/http/dtos.py index 62dea59f..d9a605a9 100644 --- a/services/api/src/api/adapters/http/dtos.py +++ b/services/api/src/api/adapters/http/dtos.py @@ -21,6 +21,9 @@ class CreateAS2TradingPartnerRequest(BaseModel): private_key_vault_ref: str | None = Field( None, max_length=512, description="Vault reference for private key (Local only)" ) + private_key_pem: str | None = Field( + None, description="Raw private key PEM to store in Vault on creation (Local only)" + ) class CreateAS2PartnershipRequest(BaseModel): diff --git a/services/api/src/api/adapters/transaction_repository.py b/services/api/src/api/adapters/transaction_repository.py index 8f548df2..f9c09030 100644 --- a/services/api/src/api/adapters/transaction_repository.py +++ b/services/api/src/api/adapters/transaction_repository.py @@ -184,18 +184,30 @@ def _apply_dynamic_filters(self, stmt: Any, model: Any, filters: list[dict[str, if field.startswith("business_metadata.") and hasattr(model, "business_metadata"): json_key = field.split("business_metadata.")[1] - column = model.business_metadata[json_key].astext + column = model.business_metadata[json_key] + column_astext = column.astext if operator == "eq": - stmt = stmt.where(column == str(value)) + from sqlalchemy import or_ + + stmt = stmt.where(or_(column_astext == str(value), column.contains(value))) elif operator == "neq": - stmt = stmt.where(column != str(value)) + from sqlalchemy import and_ + + stmt = stmt.where(and_(column_astext != str(value), ~column.contains(value))) elif operator == "contains": escaped_value = ( str(value).replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") ) - stmt = stmt.where(column.ilike(f"%{escaped_value}%", escape="\\")) + stmt = stmt.where(column_astext.ilike(f"%{escaped_value}%", escape="\\")) elif operator == "in" and isinstance(value, list): - stmt = stmt.where(column.in_([str(v) for v in value])) + from sqlalchemy import or_ + + stmt = stmt.where( + or_( + column_astext.in_([str(v) for v in value]), + *[column.contains(v) for v in value], + ) + ) continue if not hasattr(model, field): diff --git a/services/api/src/api/routers/platform/__init__.py b/services/api/src/api/routers/platform/__init__.py index 703a3308..e19da731 100644 --- a/services/api/src/api/routers/platform/__init__.py +++ b/services/api/src/api/routers/platform/__init__.py @@ -1,8 +1,10 @@ -from fastapi import APIRouter +from fastapi import APIRouter, Depends + +from api.dependencies import require_platform_admin from .scheduler import router as scheduler_router -router = APIRouter(prefix="/api/v1/platform") +router = APIRouter(prefix="/api/v1/platform", dependencies=[Depends(require_platform_admin)]) router.include_router(scheduler_router) __all__ = ["router"] diff --git a/services/api/src/api/routers/platform/scheduler.py b/services/api/src/api/routers/platform/scheduler.py index f69cab94..4ae81c5e 100644 --- a/services/api/src/api/routers/platform/scheduler.py +++ b/services/api/src/api/routers/platform/scheduler.py @@ -17,9 +17,16 @@ class JobResponse(BaseModel): id: uuid.UUID name: str status: str + target_queue: str | None + app_namespace: str | None + cron_expression: str | None + timezone: str | None next_run_at: datetime | None locked_at: datetime | None locked_by: str | None + interval_seconds: int | None + min_interval_seconds: int | None = None + max_interval_seconds: int | None = None retry_count: int error_message: str | None created_at: datetime @@ -28,6 +35,19 @@ class JobResponse(BaseModel): model_config = ConfigDict(from_attributes=True) +class JobCreateRequest(BaseModel): + name: str + interval_seconds: int = 60 + payload: dict[str, Any] = {} + + +class JobUpdateRequest(BaseModel): + interval_seconds: int | None = None + cron_expression: str | None = None + timezone: str | None = None + status: str | None = None + + class ConfigUpdateRequest(BaseModel): value: dict[str, Any] | list[Any] | str | int | bool | None @@ -47,6 +67,141 @@ async def list_jobs(uow: UnitOfWork = Depends(get_uow)) -> list[JobResponse]: return [JobResponse.model_validate(job) for job in jobs] +@router.post("/jobs", response_model=JobResponse) +async def create_job(request: JobCreateRequest, uow: UnitOfWork = Depends(get_uow)) -> JobResponse: + """Create a new scheduled background job.""" + from datetime import UTC + + from fastapi import HTTPException + from scheduler.domain.models import JobStatus + from sqlalchemy.exc import IntegrityError + + now = datetime.now(UTC) + async with uow: + try: + async with uow.global_session.begin_nested(): + new_job = ScheduledJob( + id=uuid.uuid4(), + name=request.name, + payload=request.payload, + interval_seconds=request.interval_seconds, + status=JobStatus.PENDING.value, + next_run_at=now, + retry_count=0, + max_retries=3, + created_at=now, + updated_at=now, + ) + uow.global_session.add(new_job) + await uow.global_session.flush() + except IntegrityError as e: + raise HTTPException( + status_code=409, detail=f"Job '{request.name}' already exists" + ) from e + + await uow.commit() + return JobResponse.model_validate(new_job) + + +@router.put("/jobs/{name}", response_model=JobResponse) +async def update_job( + name: str, request: JobUpdateRequest, uow: UnitOfWork = Depends(get_uow) +) -> JobResponse: + """Update a specific scheduled job.""" + from fastapi import HTTPException + from sqlalchemy import select + + async with uow: + stmt = select(ScheduledJob).where(ScheduledJob.name == name) + result = await uow.global_session.execute(stmt) + job = result.scalar_one_or_none() + if not job: + raise HTTPException(status_code=404, detail=f"Job '{name}' not found") + + if request.interval_seconds is not None: + if request.interval_seconds <= 0: + raise HTTPException(status_code=422, detail="Interval must be positive") + if ( + job.min_interval_seconds is not None + and request.interval_seconds < job.min_interval_seconds + ): + raise HTTPException(status_code=422, detail="Interval too low") + if ( + job.max_interval_seconds is not None + and request.interval_seconds > job.max_interval_seconds + ): + raise HTTPException(status_code=422, detail="Interval too high") + + job.interval_seconds = request.interval_seconds + job.cron_expression = None + job.timezone = None + + if request.cron_expression is not None: + from croniter import croniter # type: ignore + + if not croniter.is_valid(request.cron_expression): + raise HTTPException(status_code=422, detail="Invalid cron expression") + + job.cron_expression = request.cron_expression + if request.timezone is not None: + import zoneinfo + + try: + zoneinfo.ZoneInfo(request.timezone) + job.timezone = request.timezone + except zoneinfo.ZoneInfoNotFoundError as e: + raise HTTPException(status_code=422, detail="Invalid timezone") from e + + job.interval_seconds = None + + if request.status is not None: + from datetime import UTC, datetime + + from scheduler.domain.models import JobStatus + + if request.status not in [s.value for s in JobStatus]: + raise HTTPException(status_code=422, detail="Invalid status") + + if job.status == JobStatus.PAUSED.value and request.status == JobStatus.PENDING.value: + job.next_run_at = datetime.now(UTC) + + job.status = request.status + + import uuid + + from database.models.control_plane import SystemAuditLog + + audit_log = SystemAuditLog( + trace_id=uuid.uuid4(), + tenant_id=0, + event=f"SCHEDULER_JOB_UPDATED: {name}", + status="SUCCESS", + ) + uow.global_session.add(audit_log) + + await uow.global_session.flush() + await uow.commit() + + return JobResponse.model_validate(job) + + +@router.delete("/jobs/{name}", status_code=204) +async def delete_job(name: str, uow: UnitOfWork = Depends(get_uow)) -> None: + """Delete a scheduled background job.""" + from fastapi import HTTPException + from sqlalchemy import select + + async with uow: + stmt = select(ScheduledJob).where(ScheduledJob.name == name) + result = await uow.global_session.execute(stmt) + job = result.scalar_one_or_none() + if not job: + raise HTTPException(status_code=404, detail=f"Job '{name}' not found") + + await uow.global_session.delete(job) + await uow.commit() + + @router.get("/config", response_model=list[ConfigResponse]) async def get_all_config(uow: UnitOfWork = Depends(get_uow)) -> list[ConfigResponse]: """Get all platform configuration values.""" @@ -72,56 +227,7 @@ async def update_config( key: str, request: ConfigUpdateRequest, uow: UnitOfWork = Depends(get_uow) ) -> ConfigResponse: """Update a platform configuration value.""" - import datetime - import uuid - - from database.models.scheduled_job import ScheduledJob - async with uow: await uow.platform_settings.set_config(key, request.value) - - # Event-driven scheduler integration - if key == "outbox_sweeper_enabled": - stmt = select(ScheduledJob).where(ScheduledJob.name == "outbox_sweeper") - result = await uow.global_session.execute(stmt) - job = result.scalar_one_or_none() - now = datetime.datetime.now(datetime.UTC) - - if request.value is True: - if job: - job.next_run_at = now - job.status = "PENDING" - else: - # Get interval if exists - interval_cfg = await uow.platform_settings.get_config( - "outbox_sweeper_interval_seconds" - ) - interval = interval_cfg if interval_cfg is not None else 60 - - new_job = ScheduledJob( - id=uuid.uuid4(), - name="outbox_sweeper", - payload={"interval_seconds": int(interval)}, - status="PENDING", - next_run_at=now, - retry_count=0, - max_retries=3, - created_at=now, - updated_at=now, - ) - uow.global_session.add(new_job) - else: - if job: - await uow.global_session.delete(job) - - elif key == "outbox_sweeper_interval_seconds": - stmt = select(ScheduledJob).where(ScheduledJob.name == "outbox_sweeper") - result = await uow.global_session.execute(stmt) - job = result.scalar_one_or_none() - if job: - payload = dict(job.payload) if job.payload else {} - payload["interval_seconds"] = int(str(request.value)) - job.payload = payload - await uow.commit() return ConfigResponse(key=key, value=request.value) 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 bf4ca5ca..9f667455 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 @@ -60,8 +60,25 @@ async def generate_certificate( async def delete_certificate_secret( vault_ref: str, vault: VaultPort = Depends(get_vault), + uow: UnitOfWork = Depends(get_uow), ) -> None: """Deletes an orphaned private key from Vault if the UI discards it before saving.""" + async with uow: + from database.models.control_plane import AS2Partner + from sqlalchemy import or_, select + + stmt = select(AS2Partner).where( + or_( + AS2Partner.private_key_vault_ref == vault_ref, + AS2Partner.prev_private_key_vault_ref == vault_ref, + ) + ) + res = await uow.global_session.execute(stmt) + if res.scalars().first() is not None: + raise HTTPException( + status_code=400, detail="Cannot delete a private key that is currently in use." + ) + vault.delete_secret(vault_ref) @@ -85,20 +102,29 @@ async def create_platform_as2_partner( private_key_vault_ref = request.private_key_vault_ref auto_generated = False - if request.is_local and not private_key_vault_ref: - auto_generated = True - # Auto-generate self-signed cert if not already generated and provided - private_key_bytes, public_cert_bytes = generate_self_signed_cert( - common_name=request.as2_id - ) - - # Store in Vault - private_key_vault_ref = vault.store_private_key( - private_key_pem=private_key_bytes, - alias_prefix=request.name.replace(" ", "_").lower(), - ) - public_cert_pem = public_cert_bytes.decode("utf-8") + if request.is_local: + if private_key_vault_ref: + # Pre-stored vault ref (from generate cert flow) — use as-is + pass + elif request.private_key_pem: + # User uploaded their own cert+key — store the private key in Vault + auto_generated = True + private_key_vault_ref = vault.store_private_key( + private_key_pem=request.private_key_pem.encode(), + alias_prefix=request.name.replace(" ", "_").lower(), + ) + else: + # No cert material provided at all — auto-generate a self-signed cert + auto_generated = True + private_key_bytes, public_cert_bytes = generate_self_signed_cert( + common_name=request.as2_id + ) + private_key_vault_ref = vault.store_private_key( + private_key_pem=private_key_bytes, + alias_prefix=request.name.replace(" ", "_").lower(), + ) + public_cert_pem = public_cert_bytes.decode("utf-8") cmd = CreateAS2TradingPartnerCmd( name=request.name, @@ -127,10 +153,14 @@ async def create_platform_as2_partner( url=p.url, active=p.active, ) - except IntegrityError as e: + except Exception as e: if auto_generated and private_key_vault_ref: vault.delete_secret(private_key_vault_ref) - raise HTTPException(status_code=400, detail="AS2 ID already exists for this tenant.") from e + if isinstance(e, IntegrityError): + raise HTTPException( + status_code=400, detail="AS2 ID already exists for this tenant." + ) from e + raise @router.get("/as2/trading-partners", response_model=list[AS2TradingPartnerResponse]) diff --git a/services/api/src/api/services/api_receiver_service.py b/services/api/src/api/services/api_receiver_service.py index 62639a0d..be157be7 100644 --- a/services/api/src/api/services/api_receiver_service.py +++ b/services/api/src/api/services/api_receiver_service.py @@ -59,13 +59,27 @@ async def process_api_edi_json( if st: transaction_type = st.get("ST01") - business_metadata = {} + business_metadata: dict[str, Any] = {} 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]) + extracted_list = [ + self.extractor.extract(transaction_type or "", item) for item in payload + ] + for extracted in extracted_list: + for k, v in extracted.items(): + if k not in business_metadata: + business_metadata[k] = [] + # Avoid duplicates + if v not in business_metadata[k]: + business_metadata[k].append(v) - business_metadata["_routing"] = {"trading_partner_id": trading_partner_id} # type: ignore + # Flatten single-item lists for backward compatibility + for k, v in business_metadata.items(): + if isinstance(v, list) and len(v) == 1: + business_metadata[k] = v[0] + + business_metadata["_routing"] = {"trading_partner_id": trading_partner_id} # 2. Create Trace ID trace_id = uuid.uuid4() diff --git a/services/api/tests/test_scheduler.py b/services/api/tests/test_scheduler.py new file mode 100644 index 00000000..7ae24b82 --- /dev/null +++ b/services/api/tests/test_scheduler.py @@ -0,0 +1,130 @@ +from unittest.mock import AsyncMock, MagicMock + +import pytest +from api.dependencies import get_uow, require_platform_admin +from api.main import app +from fastapi.testclient import TestClient +from scheduler.domain.models import JobName + + +@pytest.fixture +def mock_uow(): + uow = AsyncMock() + # Support async with + uow.__aenter__.return_value = uow + uow.__aexit__.return_value = None + + # Mock settings get/set + uow.platform_settings = AsyncMock() + uow.platform_settings.get_config.return_value = 60 + + # Mock global session + uow.global_session = AsyncMock() + uow.global_session.begin_nested = MagicMock() + uow.global_session.begin_nested.return_value.__aenter__ = AsyncMock() + uow.global_session.begin_nested.return_value.__aexit__ = AsyncMock(return_value=None) + + # Setup execute return + mock_result = MagicMock() + mock_result.scalars().all.return_value = [] + mock_result.scalar_one_or_none.return_value = None + uow.global_session.execute.return_value = mock_result + + return uow + + +@pytest.fixture +def client(mock_uow): + app.dependency_overrides[get_uow] = lambda: mock_uow + app.dependency_overrides[require_platform_admin] = lambda: 0 + yield TestClient(app) + app.dependency_overrides.clear() + + +def test_scheduler_list_jobs_empty(client): + resp = client.get("/api/v1/platform/scheduler/jobs") + assert resp.status_code == 200 + assert resp.json() == [] + + +def test_scheduler_get_all_config(client): + resp = client.get("/api/v1/platform/scheduler/config") + assert resp.status_code == 200 + assert isinstance(resp.json(), list) + + +def test_scheduler_get_config(client): + resp = client.get("/api/v1/platform/scheduler/config/some_config_key") + assert resp.status_code == 200 + assert resp.json()["key"] == "some_config_key" + + +def test_scheduler_update_job_invalid_interval(client, mock_uow): + import uuid + from datetime import datetime + + mock_job = MagicMock() + mock_job.id = uuid.uuid4() + mock_job.name = JobName.OUTBOX_SWEEPER + mock_job.status = "PENDING" + mock_job.next_run_at = datetime.now() + mock_job.locked_at = None + mock_job.locked_by = None + mock_job.interval_seconds = 60 + mock_job.min_interval_seconds = 10 + mock_job.max_interval_seconds = 300 + mock_job.retry_count = 0 + mock_job.error_message = None + mock_job.created_at = datetime.now() + mock_job.updated_at = datetime.now() + mock_result = MagicMock() + mock_result.scalar_one_or_none.return_value = mock_job + mock_uow.global_session.execute.return_value = mock_result + + resp = client.put( + "/api/v1/platform/scheduler/jobs/outbox_sweeper", + json={"interval_seconds": -5}, + ) + assert resp.status_code == 422 + + +def test_scheduler_create_job(client, mock_uow): + resp = client.post( + "/api/v1/platform/scheduler/jobs", + json={"name": "test_job", "interval_seconds": 60, "payload": {}}, + ) + assert resp.status_code == 200 + assert mock_uow.global_session.add.called + assert resp.json()["name"] == "test_job" + + +def test_scheduler_create_job_already_exists(client, mock_uow): + from sqlalchemy.exc import IntegrityError + + mock_uow.global_session.flush.side_effect = IntegrityError("mock", {}, Exception()) + + resp = client.post( + "/api/v1/platform/scheduler/jobs", + json={"name": "test_job"}, + ) + assert resp.status_code == 409 + + +def test_scheduler_delete_job(client, mock_uow): + mock_job = MagicMock() + mock_result = MagicMock() + mock_result.scalar_one_or_none.return_value = mock_job + mock_uow.global_session.execute.return_value = mock_result + + resp = client.delete("/api/v1/platform/scheduler/jobs/test_job") + assert resp.status_code == 204 + assert mock_uow.global_session.delete.called + + +def test_scheduler_delete_job_not_found(client, mock_uow): + mock_result = MagicMock() + mock_result.scalar_one_or_none.return_value = None + mock_uow.global_session.execute.return_value = mock_result + + resp = client.delete("/api/v1/platform/scheduler/jobs/test_job") + assert resp.status_code == 404 diff --git a/services/as2_server/scripts/seed.py b/services/as2_server/scripts/seed.py index 16f47f4c..5642a854 100644 --- a/services/as2_server/scripts/seed.py +++ b/services/as2_server/scripts/seed.py @@ -97,6 +97,52 @@ async def seed_database() -> None: else: logger.info("SYSTEM_ADMIN_EMAIL not provided. Skipping default admin creation.") + # 4. Seed Core System Jobs + logger.info("Seeding Core System Jobs...") + from datetime import UTC, datetime + + from database.models.scheduled_job import ScheduledJob + from scheduler.domain.models import JobStatus + from scheduler.registry import SYSTEM_JOB_REGISTRY + + now = datetime.now(UTC) + for job_def in SYSTEM_JOB_REGISTRY: + job_result = await session.execute(select(ScheduledJob).filter_by(name=job_def.name)) + job = job_result.scalar_one_or_none() + + if not job: + job = ScheduledJob( + name=job_def.name, + payload={}, + status=JobStatus.PENDING.value, + target_queue=job_def.target_queue, + app_namespace=job_def.app_namespace, + interval_seconds=job_def.default_interval_seconds, + cron_expression=job_def.default_cron_expression, + timezone=job_def.default_timezone, + min_interval_seconds=job_def.min_interval_seconds, + max_interval_seconds=job_def.max_interval_seconds, + retry_count=0, + max_retries=job_def.max_retries, + created_at=now, + updated_at=now, + next_run_at=now, + ) + session.add(job) + logger.info(f"Created system job: {job_def.name}.") + else: + # Always sync canonical config bounds from the registry, + # ensuring existing rows stay consistent after schema migrations. + job.target_queue = job_def.target_queue + job.app_namespace = job_def.app_namespace + job.cron_expression = job_def.default_cron_expression + job.timezone = job_def.default_timezone + job.min_interval_seconds = job_def.min_interval_seconds + job.max_interval_seconds = job_def.max_interval_seconds + logger.info(f"Synced config for system job: {job_def.name}.") + + await session.flush() + await session.commit() logger.info("Database seed completed successfully.") diff --git a/services/workers/orchestrator/src/worker/adapters/sqs_poller.py b/services/workers/orchestrator/src/worker/adapters/sqs_poller.py new file mode 100644 index 00000000..a7c4d781 --- /dev/null +++ b/services/workers/orchestrator/src/worker/adapters/sqs_poller.py @@ -0,0 +1,87 @@ +import asyncio +import json +import logging +from collections.abc import Callable +from typing import Any + +import aioboto3 # type: ignore[import-untyped] + +logger = logging.getLogger(__name__) + + +async def poll_sqs_queue( + queue_name: str, + processor_func: Callable[[dict[str, Any]], Any], + aws_endpoint: str | None = None, +) -> None: + """Long-polls an SQS queue and processes messages.""" + session = aioboto3.Session() + client_kwargs = {"region_name": "us-east-1"} + if aws_endpoint: + client_kwargs["endpoint_url"] = aws_endpoint + + while True: + try: + async with session.client("sqs", **client_kwargs) as sqs: + queue_url_resp = await sqs.get_queue_url(QueueName=queue_name) + queue_url = queue_url_resp["QueueUrl"] + + logger.info(f"Started polling {queue_name} ({queue_url})") + + while True: + response = await sqs.receive_message( + QueueUrl=queue_url, + MaxNumberOfMessages=10, + WaitTimeSeconds=20, + ) + + messages = response.get("Messages", []) + for msg in messages: + receipt_handle = msg["ReceiptHandle"] + try: + body = json.loads(msg["Body"]) + + # Log trace_id if available, otherwise just log processing + trace_id = body.get("payload", {}).get("trace_id") + if trace_id: + logger.info(f"[{queue_name}] Processing trace_id={trace_id}") + else: + logger.info(f"[{queue_name}] Processing message") + + await processor_func(body) + + # Delete message on success + await sqs.delete_message( + QueueUrl=queue_url, ReceiptHandle=receipt_handle + ) + + if trace_id: + logger.info( + f"[{queue_name}] Successfully processed trace_id={trace_id}" + ) + else: + logger.info(f"[{queue_name}] Successfully processed message") + + except json.JSONDecodeError: + # Permanently delete malformed (non-JSON) messages + logger.error( + f"[{queue_name}] Non-JSON message body, deleting permanently" + ) + await sqs.delete_message( + QueueUrl=queue_url, ReceiptHandle=receipt_handle + ) + except (KeyError, NotImplementedError) as e: + # Permanently delete messages with unrecoverable configuration errors + logger.error( + f"[{queue_name}] Permanent validation error, deleting permanently: {e}" + ) + await sqs.delete_message( + QueueUrl=queue_url, ReceiptHandle=receipt_handle + ) + except Exception as e: + logger.exception( + f"[{queue_name}] Transient error processing message: {e}" + ) + except Exception as e: + logger.exception(f"[{queue_name}] SQS client error, retrying in 2s: {e}") + await asyncio.sleep(2) diff --git a/services/workers/orchestrator/src/worker/adapters/sqs_publisher.py b/services/workers/orchestrator/src/worker/adapters/sqs_publisher.py new file mode 100644 index 00000000..a741c13b --- /dev/null +++ b/services/workers/orchestrator/src/worker/adapters/sqs_publisher.py @@ -0,0 +1,110 @@ +import json +import logging +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager +from typing import Any + +import aioboto3 # type: ignore[import-untyped] + +from worker.ports.message_publisher import MessagePublisherPort + +logger = logging.getLogger(__name__) + + +class SqsPublisherAdapter(MessagePublisherPort): + def __init__(self, endpoint_url: str | None = None, region: str = "us-east-1"): + self.endpoint_url = endpoint_url + self.region = region + self.session = aioboto3.Session() + self._queue_url_cache: dict[str, str] = {} + self._sqs_client = None + + @asynccontextmanager + async def connect(self) -> AsyncIterator["MessagePublisherPort"]: + async with self.session.client( + "sqs", endpoint_url=self.endpoint_url, region_name=self.region + ) as sqs: + self._sqs_client = sqs + try: + yield self + finally: + self._sqs_client = None + + async def publish_batch(self, queue_name: str, messages: list[dict[str, Any]]) -> list[str]: + if not messages: + return [] + + if self._sqs_client is None: + raise RuntimeError("publish_batch must be called within the connect() context manager") + + sqs = self._sqs_client + successful_ids = [] + + if queue_name not in self._queue_url_cache: + try: + resp = await sqs.get_queue_url(QueueName=queue_name) + self._queue_url_cache[queue_name] = resp["QueueUrl"] + except Exception: + logger.exception(f"Failed to get queue url for {queue_name}") + return [] + + queue_url = self._queue_url_cache[queue_name] + + # SQS allows max 10 messages per batch + for i in range(0, len(messages), 10): + batch = messages[i : i + 10] + entries = [] + for msg in batch: + # 'Id' is required and used to correlate success/failure responses. + # We pass the full message dict (which must contain 'Id'). + if "Id" not in msg: + logger.warning("Message missing 'Id' key, skipping.") + continue + + entries.append( + { + "Id": msg["Id"], + "MessageBody": json.dumps(msg["MessageBody"]), + } + ) + + if not entries: + continue + + try: + resp = await sqs.send_message_batch(QueueUrl=queue_url, Entries=entries) + + for success in resp.get("Successful", []): + successful_ids.append(success["Id"]) + + for failed in resp.get("Failed", []): + logger.error( + f"Failed to forward message id={failed['Id']}: {failed['Message']}" + ) + except Exception: + logger.exception(f"Failed to send batch to {queue_name}") + + return successful_ids + + async def publish(self, queue_name: str, payload: dict) -> None: + """Publishes a single message without requiring the connect context manager.""" + if queue_name not in self._queue_url_cache: + try: + async with self.session.client( + "sqs", endpoint_url=self.endpoint_url, region_name=self.region + ) as sqs: + resp = await sqs.get_queue_url(QueueName=queue_name) + self._queue_url_cache[queue_name] = resp["QueueUrl"] + except Exception: + logger.exception(f"Failed to get queue url for {queue_name}") + return + + queue_url = self._queue_url_cache[queue_name] + + try: + async with self.session.client( + "sqs", endpoint_url=self.endpoint_url, region_name=self.region + ) as sqs: + await sqs.send_message(QueueUrl=queue_url, MessageBody=json.dumps(payload)) + except Exception: + logger.exception(f"Failed to send single message to {queue_name}") diff --git a/services/workers/orchestrator/src/worker/core/job_registry.py b/services/workers/orchestrator/src/worker/core/job_registry.py new file mode 100644 index 00000000..783b8591 --- /dev/null +++ b/services/workers/orchestrator/src/worker/core/job_registry.py @@ -0,0 +1,16 @@ +import logging +from typing import Any + +logger = logging.getLogger(__name__) + + +class JobHandlerRegistry: + def __init__(self) -> None: + self._handlers: dict[str, Any] = {} + + def register(self, job_name: str, handler: Any) -> None: + self._handlers[job_name] = handler + logger.info(f"Registered job handler for {job_name}") + + def get(self, job_name: str) -> Any | None: + return self._handlers.get(job_name) diff --git a/services/workers/orchestrator/src/worker/core/security.py b/services/workers/orchestrator/src/worker/core/security.py new file mode 100644 index 00000000..5dd9f687 --- /dev/null +++ b/services/workers/orchestrator/src/worker/core/security.py @@ -0,0 +1,52 @@ +import ipaddress +import logging +from urllib.parse import urlparse + +logger = logging.getLogger(__name__) + + +def validate_target_url(url: str) -> bool: + """ + Validate target URL to prevent SSRF attacks. + Returns True if URL is safe, False otherwise. + """ + try: + parsed = urlparse(url) + + # Only allow http and https schemes + if parsed.scheme not in ("http", "https"): + logger.warning(f"SSRF check failed: invalid scheme {parsed.scheme}") + return False + + # Reject URLs without a hostname + if not parsed.hostname: + logger.warning("SSRF check failed: missing hostname") + return False + + # Resolve all A/AAAA records for the hostname + import socket + + try: + # getaddrinfo returns a list of 5-tuples: (family, type, proto, canonname, sockaddr) + addr_info = socket.getaddrinfo(parsed.hostname, None) + except socket.gaierror: + logger.warning(f"SSRF check failed: could not resolve hostname {parsed.hostname}") + return False + + for addr in addr_info: + ip_str = addr[4][0] + ip = ipaddress.ip_address(ip_str) + if ( + ip.is_private + or ip.is_loopback + or ip.is_link_local + or ip.is_reserved + or ip.is_multicast + ): + logger.warning(f"SSRF check failed: resolved to private/internal IP {ip}") + return False + + return True + except Exception as e: + logger.error(f"SSRF validation error: {e}") + return False diff --git a/services/workers/orchestrator/src/worker/core/tenant_resolver.py b/services/workers/orchestrator/src/worker/core/tenant_resolver.py new file mode 100644 index 00000000..6d4be95a --- /dev/null +++ b/services/workers/orchestrator/src/worker/core/tenant_resolver.py @@ -0,0 +1,37 @@ +from database.connection import DatabaseRouter +from database.models.control_plane import DatabaseShard, Tenant +from sqlalchemy import select + + +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() diff --git a/services/workers/orchestrator/src/worker/data/handlers.py b/services/workers/orchestrator/src/worker/data/handlers.py new file mode 100644 index 00000000..da829b36 --- /dev/null +++ b/services/workers/orchestrator/src/worker/data/handlers.py @@ -0,0 +1,197 @@ +import contextlib +import logging +from typing import Any + +from config.settings import get_settings +from database.connection import DatabaseRouter +from pipeline.adapters.as2 import HttpxAS2DeliveryAdapter +from pipeline.adapters.http import HttpxDeliveryAdapter +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.delivery import ( + As2DeliveryStrategy, + DeliveryRouter, + SftpDeliveryStrategy, + WebhookDeliveryStrategy, +) +from pipeline.core.transformation import InboundTransformService, OutboundTransformService +from worker.adapters.vault import WorkerVaultAdapter +from worker.core.security import validate_target_url +from worker.core.tenant_resolver import TenantResolver + +logger = logging.getLogger(__name__) + + +async def process_pipeline_event( + trace_id: str, + event_type: str, + payload: dict[str, Any], + tenant_id: int, + resolver: TenantResolver, + db_router: DatabaseRouter, + s3_bucket: str, + aws_endpoint: str | None, + idempotency_key: str | None = None, +) -> None: + """Sets up the Hexagonal dependencies and executes TransformService or Saga Coordinator.""" + shard_name, shard_dsn = await resolver.resolve(tenant_id) + + tenant_gen = db_router.get_tenant_session(tenant_id, shard_name, shard_dsn) + session = await tenant_gen.__anext__() + try: + 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() + + if idempotency_key: + import uuid + + from database.models.data_plane import DataPlaneOutbox, ProcessedEvent + from sqlalchemy import select, update + + key_uuid = uuid.UUID(idempotency_key) + + # Check for duplicate + stmt = select(ProcessedEvent).where(ProcessedEvent.idempotency_key == key_uuid) + existing = await session.execute(stmt) + if existing.scalar_one_or_none(): + logger.info(f"Skipping duplicate event with idempotency_key={idempotency_key}") + await session.commit() + return + + # Mark as processed in same transaction + session.add(ProcessedEvent(idempotency_key=key_uuid)) + await session.execute( + update(DataPlaneOutbox) + .where(DataPlaneOutbox.idempotency_key == key_uuid) + .values(status="PROCESSED") + ) + + from domain.events import PipelineEventType + + if event_type in ( + PipelineEventType.TRANSFORM_COMPLETED, + PipelineEventType.DELIVERY_COMPLETED, + ): + from pipeline.core.saga import TraceLifecycleService + + saga_service = TraceLifecycleService(repo_adapter) + if event_type == PipelineEventType.TRANSFORM_COMPLETED: + await saga_service.handle_transform_completed(payload) + else: + await saga_service.handle_delivery_completed(payload) + else: + # Resolve direction first + from domain.direction import MessageDirection + + direction_str = payload.get("direction", MessageDirection.INBOUND.value) + direction = ( + MessageDirection.OUTBOUND + if direction_str.upper() == MessageDirection.OUTBOUND.value + else MessageDirection.INBOUND + ) + + # Execute pure domain logic + service = ( + InboundTransformService(transformer_adapter, repo_adapter) + if direction == MessageDirection.INBOUND + else OutboundTransformService(transformer_adapter, repo_adapter) + ) + print(f"[WORKER] Transforming trace_id={trace_id}") + + await service.transform(trace_id) + print(f"[WORKER] SUCCESS transforming trace_id={trace_id}") + + # Commit transaction + await session.commit() + except Exception as e: + print(f"[WORKER] FAILURE in process_transformation for trace_id={trace_id}: {e}") + await session.rollback() + raise + finally: + with contextlib.suppress(StopAsyncIteration): + await tenant_gen.__anext__() + + +async def process_delivery( + trace_id: str, + event_type: str, + payload: dict[str, Any], + tenant_id: int, + resolver: TenantResolver, + db_router: DatabaseRouter, + s3_bucket: str, + aws_endpoint: str | None, + idempotency_key: str | None = None, +) -> None: + """Sets up the Hexagonal dependencies and executes DeliveryService.""" + shard_name, shard_dsn = await resolver.resolve(tenant_id) + + tenant_gen = db_router.get_tenant_session(tenant_id, shard_name, shard_dsn) + session = await tenant_gen.__anext__() + try: + 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() + + if idempotency_key: + import uuid + + from database.models.data_plane import DataPlaneOutbox, ProcessedEvent + from sqlalchemy import select, update + + key_uuid = uuid.UUID(idempotency_key) + + # Check for duplicate + stmt = select(ProcessedEvent).where(ProcessedEvent.idempotency_key == key_uuid) + existing = await session.execute(stmt) + if existing.scalar_one_or_none(): + logger.info( + f"Skipping duplicate delivery event with idempotency_key={idempotency_key}" + ) + await session.commit() + return + + # Mark as processed in same transaction + session.add(ProcessedEvent(idempotency_key=key_uuid)) + await session.execute( + update(DataPlaneOutbox) + .where(DataPlaneOutbox.idempotency_key == key_uuid) + .values(status="PROCESSED") + ) + + # Instantiate Domain Service + strategies = { + "webhook_id": WebhookDeliveryStrategy(repo_adapter, http_adapter, vault_adapter), + "sftp_partner_id": SftpDeliveryStrategy(repo_adapter, sftp_adapter, vault_adapter), + "as2_partner_id": As2DeliveryStrategy(repo_adapter, as2_adapter, vault_adapter), + } + service = DeliveryRouter( + repository=repo_adapter, + strategies=strategies, + ) + + # Execute pure domain logic + await service.deliver(trace_id) + + # Commit transaction + await session.commit() + except Exception: + await session.rollback() + raise + finally: + with contextlib.suppress(StopAsyncIteration): + await tenant_gen.__anext__() diff --git a/services/workers/orchestrator/src/worker/data/main.py b/services/workers/orchestrator/src/worker/data/main.py index e5784a27..d25fab5d 100644 --- a/services/workers/orchestrator/src/worker/data/main.py +++ b/services/workers/orchestrator/src/worker/data/main.py @@ -1,341 +1,23 @@ import asyncio -import contextlib -import ipaddress -import json import logging import os -from collections.abc import Callable from typing import Any -from urllib.parse import urlparse -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.repository import SqlAlchemyRepositoryAdapter -from pipeline.adapters.sftp import ParamikoSftpDeliveryAdapter -from pipeline.adapters.storage import S3StorageAdapter -from pipeline.adapters.transformer import BotsTransformerAdapter -from pipeline.core.delivery import ( - As2DeliveryStrategy, - DeliveryRouter, - SftpDeliveryStrategy, - WebhookDeliveryStrategy, -) -from pipeline.core.transformation import InboundTransformService, OutboundTransformService from scheduler.adapters.repository import SqlAlchemyJobRepository from scheduler.core.service import SchedulerWorkerService -from sqlalchemy import select -from worker.adapters.vault import WorkerVaultAdapter -from worker.jobs.outbox_sweeper import DataPlaneOutboxSweeperJobHandler +from worker.adapters.sqs_poller import poll_sqs_queue +from worker.adapters.sqs_publisher import SqsPublisherAdapter +from worker.core.tenant_resolver import TenantResolver +from worker.data.handlers import process_delivery, process_pipeline_event 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__) -def validate_target_url(url: str) -> bool: - """ - Validate target URL to prevent SSRF attacks. - Returns True if URL is safe, False otherwise. - """ - try: - parsed = urlparse(url) - - # Only allow http and https schemes - if parsed.scheme not in ("http", "https"): - logger.warning(f"SSRF check failed: invalid scheme {parsed.scheme}") - return False - - # Reject URLs without a hostname - if not parsed.hostname: - logger.warning("SSRF check failed: missing hostname") - return False - - # Resolve all A/AAAA records for the hostname - import socket - - try: - # getaddrinfo returns a list of 5-tuples: (family, type, proto, canonname, sockaddr) - addr_info = socket.getaddrinfo(parsed.hostname, None) - except socket.gaierror: - logger.warning(f"SSRF check failed: could not resolve hostname {parsed.hostname}") - return False - - for addr in addr_info: - ip_str = addr[4][0] - ip = ipaddress.ip_address(ip_str) - if ( - ip.is_private - or ip.is_loopback - or ip.is_link_local - or ip.is_reserved - or ip.is_multicast - ): - logger.warning(f"SSRF check failed: resolved to private/internal IP {ip}") - return False - - return True - except Exception as e: - logger.error(f"SSRF validation error: {e}") - return False - - -async def process_pipeline_event( - trace_id: str, - event_type: str, - payload: dict[str, Any], - tenant_id: int, - resolver: TenantResolver, - db_router: DatabaseRouter, - s3_bucket: str, - aws_endpoint: str | None, -) -> None: - """Sets up the Hexagonal dependencies and executes TransformService or Saga Coordinator.""" - shard_name, shard_dsn = await resolver.resolve(tenant_id) - - tenant_gen = db_router.get_tenant_session(tenant_id, shard_name, shard_dsn) - session = await tenant_gen.__anext__() - try: - 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() - - from domain.events import PipelineEventType - - if event_type in ( - PipelineEventType.TRANSFORM_COMPLETED, - PipelineEventType.DELIVERY_COMPLETED, - ): - from pipeline.core.saga import TraceLifecycleService - - saga_service = TraceLifecycleService(repo_adapter) - if event_type == PipelineEventType.TRANSFORM_COMPLETED: - await saga_service.handle_transform_completed(payload) - else: - await saga_service.handle_delivery_completed(payload) - else: - # Resolve direction first - from domain.direction import MessageDirection - - direction_str = payload.get("direction", MessageDirection.INBOUND.value) - direction = ( - MessageDirection.OUTBOUND - if direction_str.upper() == MessageDirection.OUTBOUND.value - else MessageDirection.INBOUND - ) - - # Execute pure domain logic - service = ( - InboundTransformService(transformer_adapter, repo_adapter) - if direction == MessageDirection.INBOUND - else OutboundTransformService(transformer_adapter, repo_adapter) - ) - print(f"[WORKER] Transforming trace_id={trace_id}") - - await service.transform(trace_id) - print(f"[WORKER] SUCCESS transforming trace_id={trace_id}") - - # Commit transaction - await session.commit() - except Exception as e: - print(f"[WORKER] FAILURE in process_transformation for trace_id={trace_id}: {e}") - await session.rollback() - raise - finally: - with contextlib.suppress(StopAsyncIteration): - await tenant_gen.__anext__() - - -async def process_delivery( - trace_id: str, - event_type: str, - payload: dict[str, Any], - tenant_id: int, - resolver: TenantResolver, - db_router: DatabaseRouter, - s3_bucket: str, - aws_endpoint: str | None, -) -> None: - """Sets up the Hexagonal dependencies and executes DeliveryService.""" - shard_name, shard_dsn = await resolver.resolve(tenant_id) - - tenant_gen = db_router.get_tenant_session(tenant_id, shard_name, shard_dsn) - session = await tenant_gen.__anext__() - try: - 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 - strategies = { - "webhook_id": WebhookDeliveryStrategy(repo_adapter, http_adapter, vault_adapter), - "sftp_partner_id": SftpDeliveryStrategy(repo_adapter, sftp_adapter, vault_adapter), - "as2_partner_id": As2DeliveryStrategy(repo_adapter, as2_adapter, vault_adapter), - } - service = DeliveryRouter( - repository=repo_adapter, - strategies=strategies, - ) - - # Execute pure domain logic - await service.deliver(trace_id) - - # Commit transaction - await session.commit() - except Exception: - await session.rollback() - raise - finally: - with contextlib.suppress(StopAsyncIteration): - await tenant_gen.__anext__() - - -async def poll_sqs_queue( - queue_name: str, - processor_func: Callable[..., Any], - resolver: TenantResolver, - db_router: DatabaseRouter, - s3_bucket: str, - aws_endpoint: str | None, -) -> None: - """Long-polls an SQS queue and processes messages.""" - session = aioboto3.Session() - client_kwargs = {"region_name": "us-east-1"} - if aws_endpoint: - client_kwargs["endpoint_url"] = aws_endpoint - - while True: - try: - async with session.client("sqs", **client_kwargs) as sqs: - queue_url_resp = await sqs.get_queue_url(QueueName=queue_name) - queue_url = queue_url_resp["QueueUrl"] - - logger.info(f"Started polling {queue_name} ({queue_url})") - - while True: - response = await sqs.receive_message( - QueueUrl=queue_url, - MaxNumberOfMessages=10, - WaitTimeSeconds=20, - ) - - messages = response.get("Messages", []) - for msg in messages: - receipt_handle = msg["ReceiptHandle"] - try: - body = json.loads(msg["Body"]) - payload = body.get("payload", {}) - trace_id = payload.get("trace_id") - tenant_id = body.get("tenant_id") - - if not trace_id or not tenant_id: - logger.error(f"Missing trace_id or tenant_id in message: {body}") - # Permanently delete unrecoverable messages to prevent re-drive loops - await sqs.delete_message( - QueueUrl=queue_url, ReceiptHandle=receipt_handle - ) - logger.warning( - f"[{queue_name}] Deleted poison message with missing ids" - ) - continue - - 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"), - "payload": payload, - "tenant_id": tenant_id, - "resolver": resolver, - "db_router": db_router, - "s3_bucket": s3_bucket, - "aws_endpoint": aws_endpoint, - } - await processor_func(**kwargs) - - # Delete message on success - await sqs.delete_message( - QueueUrl=queue_url, ReceiptHandle=receipt_handle - ) - logger.info( - f"[{queue_name}] Successfully processed trace_id={trace_id}" - ) - - except json.JSONDecodeError: - # Permanently delete malformed (non-JSON) messages - logger.error( - f"[{queue_name}] Non-JSON message body, deleting permanently" - ) - await sqs.delete_message( - QueueUrl=queue_url, ReceiptHandle=receipt_handle - ) - except (KeyError, NotImplementedError) as e: - # Permanently delete messages with unrecoverable configuration errors - logger.error( - f"[{queue_name}] Permanent validation error, deleting permanently: {e}" - ) - await sqs.delete_message( - QueueUrl=queue_url, ReceiptHandle=receipt_handle - ) - except Exception as e: - logger.exception( - f"[{queue_name}] Transient error processing message: {e}" - ) - except Exception as e: - logger.exception(f"[{queue_name}] SQS client error, retrying in 2s: {e}") - await asyncio.sleep(2) - - async def main() -> None: settings = get_settings() aws_endpoint = os.getenv("AWS_ENDPOINT_URL") @@ -344,23 +26,56 @@ async def main() -> None: db_router = DatabaseRouter(global_db_url=settings.database.global_url) resolver = TenantResolver(db_router) + async def pipeline_processor(body: dict[str, Any]) -> None: + payload = body.get("payload", {}) + trace_id = payload.get("trace_id") + tenant_id = body.get("tenant_id") + if not trace_id or not tenant_id: + logger.error(f"Missing trace_id or tenant_id in message: {body}") + return + await process_pipeline_event( + trace_id=trace_id, + event_type=body.get("event_type", "UNKNOWN"), + payload=payload, + tenant_id=tenant_id, + resolver=resolver, + db_router=db_router, + s3_bucket=s3_bucket, + aws_endpoint=aws_endpoint, + idempotency_key=body.get("idempotency_key"), + ) + transform_task = asyncio.create_task( poll_sqs_queue( MessageQueueName.TRANSFORM_ORCHESTRATION_QUEUE, - process_pipeline_event, - resolver, - db_router, - s3_bucket, + pipeline_processor, aws_endpoint, ) ) + + async def delivery_processor(body: dict[str, Any]) -> None: + payload = body.get("payload", {}) + trace_id = payload.get("trace_id") + tenant_id = body.get("tenant_id") + if not trace_id or not tenant_id: + logger.error(f"Missing trace_id or tenant_id in message: {body}") + return + await process_delivery( + trace_id=trace_id, + event_type=body.get("event_type", "UNKNOWN"), + payload=payload, + tenant_id=tenant_id, + resolver=resolver, + db_router=db_router, + s3_bucket=s3_bucket, + aws_endpoint=aws_endpoint, + idempotency_key=body.get("idempotency_key"), + ) + deliver_task = asyncio.create_task( poll_sqs_queue( MessageQueueName.DELIVER_QUEUE, - process_delivery, - resolver, - db_router, - s3_bucket, + delivery_processor, aws_endpoint, ) ) @@ -370,17 +85,46 @@ async def main() -> None: engine = create_async_engine(settings.database.global_url) scheduler_repo = SqlAlchemyJobRepository(engine) + message_publisher = SqsPublisherAdapter( + endpoint_url=settings.aws.endpoint_url, + region=settings.aws.resolved_region, + ) + scheduler_service = SchedulerWorkerService( - scheduler_repo, worker_id=f"orchestrator-{os.getpid()}" + scheduler_repo, publisher=message_publisher, worker_id=f"orchestrator-{os.getpid()}" + ) + + from scheduler.domain.models import JobName + from worker.core.job_registry import JobHandlerRegistry + from worker.jobs.data_retention import DataRetentionCleanupJobHandler + from worker.jobs.outbox_sweeper import DataPlaneOutboxSweeperJobHandler + + registry = JobHandlerRegistry() + registry.register( + JobName.OUTBOX_SWEEPER.value, DataPlaneOutboxSweeperJobHandler(db_router, message_publisher) + ) + registry.register( + JobName.DATA_RETENTION_CLEANUP.value, DataRetentionCleanupJobHandler(db_router) ) - scheduler_service.register_handler( - "outbox_sweeper", DataPlaneOutboxSweeperJobHandler(db_router) + + import functools + + from worker.data.scheduled_jobs_handler import process_scheduled_job + + scheduled_jobs_processor = functools.partial(process_scheduled_job, registry=registry) + + scheduled_jobs_task = asyncio.create_task( + poll_sqs_queue( + "edi-orchestrator-jobs", + scheduled_jobs_processor, + aws_endpoint, + ) ) # Run the scheduler loop in the background await scheduler_service.start(poll_interval_seconds=10.0) - await asyncio.gather(transform_task, deliver_task) + await asyncio.gather(transform_task, deliver_task, scheduled_jobs_task) if __name__ == "__main__": diff --git a/services/workers/orchestrator/src/worker/data/scheduled_jobs_handler.py b/services/workers/orchestrator/src/worker/data/scheduled_jobs_handler.py new file mode 100644 index 00000000..3936e56d --- /dev/null +++ b/services/workers/orchestrator/src/worker/data/scheduled_jobs_handler.py @@ -0,0 +1,38 @@ +import logging +import uuid +from typing import Any + +from scheduler.domain.models import Job + +logger = logging.getLogger(__name__) + + +async def process_scheduled_job(message: dict[str, Any], **kwargs: Any) -> None: + """ + Generic dispatcher for scheduled jobs. + It requires a 'registry' to be passed in via **kwargs. + """ + job_id = message.get("job_id") + job_name = message.get("job_name") + job_payload = message.get("payload", {}) + + if not job_id or not job_name: + logger.error(f"Missing job_id or job_name in message: {message}") + return + + logger.info(f"Processing scheduled job: {job_name} ({job_id})") + + registry = kwargs.get("registry") + if not registry: + logger.error("JobHandlerRegistry not found in kwargs") + return + + handler = registry.get(job_name) + if not handler: + logger.error(f"Unknown scheduled job name: {job_name}") + return + + # Reconstruct a dummy Job object just enough for the handler to execute it + job = Job(id=uuid.UUID(job_id), name=job_name, payload=job_payload) + + await handler.execute(job) diff --git a/services/workers/orchestrator/src/worker/jobs/data_retention.py b/services/workers/orchestrator/src/worker/jobs/data_retention.py new file mode 100644 index 00000000..d9b5f50f --- /dev/null +++ b/services/workers/orchestrator/src/worker/jobs/data_retention.py @@ -0,0 +1,83 @@ +import asyncio +import datetime +import logging + +from database.connection import DatabaseRouter +from database.models.control_plane import DatabaseShard +from database.models.data_plane import DataPlaneOutbox, ProcessedEvent +from scheduler.domain.models import Job +from scheduler.ports.handler import JobHandlerPort +from sqlalchemy import delete, select +from sqlalchemy.ext.asyncio import AsyncSession + +logger = logging.getLogger(__name__) + +_CONCURRENCY_LIMIT = 5 +_RETENTION_DAYS = 7 + + +class DataRetentionCleanupJobHandler(JobHandlerPort): + def __init__(self, db_router: DatabaseRouter) -> None: + self.db_router = db_router + + async def execute(self, job: Job) -> datetime.datetime | None: + """ + Cleans up old PROCESSED outbox events and processed idempotency keys + across all tenant shards to prevent unbounded database growth. + """ + logger.info(f"[DataRetentionCleanup] Running sweep for job {job.id}") + + sem = asyncio.Semaphore(_CONCURRENCY_LIMIT) + + async for global_session in self.db_router.get_global_session(): + res = await global_session.execute(select(DatabaseShard)) + shards = res.scalars().all() + + async def _bounded_cleanup(shard_name: str, shard_dsn: str) -> tuple[int, int]: + async with sem: + try: + return await self._cleanup_shard(shard_name, shard_dsn) + except Exception as e: + logger.error(f"[DataRetentionCleanup] Failed cleaning shard {shard_name}: {e}") + return 0, 0 + + results = await asyncio.gather( + *[_bounded_cleanup(shard.name, shard.dsn) for shard in shards] + ) + + total_outbox = sum(r[0] for r in results) + total_processed = sum(r[1] for r in results) + + logger.info( + f"[DataRetentionCleanup] Cleanup complete. " + f"Deleted {total_outbox} outbox rows and {total_processed} processed_events rows." + ) + + # Return None to let the scheduler calculate the next run based on interval + return None + + async def _cleanup_shard(self, shard_name: str, shard_dsn: str) -> tuple[int, int]: + """Sweep a single tenant shard, deleting old records.""" + cutoff_date = datetime.datetime.now(datetime.UTC) - datetime.timedelta(days=_RETENTION_DAYS) + + outbox_deleted = 0 + processed_deleted = 0 + + engine = await self.db_router.get_engine(shard_name, shard_dsn) + async with AsyncSession(engine, expire_on_commit=False) as session: + # Delete old PROCESSED outbox events + stmt_outbox = delete(DataPlaneOutbox).where( + DataPlaneOutbox.status == "PROCESSED", + DataPlaneOutbox.created_at < cutoff_date, + ) + res_outbox = await session.execute(stmt_outbox) + outbox_deleted = res_outbox.rowcount + + # Delete old processed_events + stmt_processed = delete(ProcessedEvent).where(ProcessedEvent.processed_at < cutoff_date) + res_processed = await session.execute(stmt_processed) + processed_deleted = res_processed.rowcount + + await session.commit() + + return outbox_deleted, processed_deleted diff --git a/services/workers/orchestrator/src/worker/jobs/outbox_sweeper.py b/services/workers/orchestrator/src/worker/jobs/outbox_sweeper.py index b7b565d4..0d831679 100644 --- a/services/workers/orchestrator/src/worker/jobs/outbox_sweeper.py +++ b/services/workers/orchestrator/src/worker/jobs/outbox_sweeper.py @@ -1,11 +1,7 @@ import asyncio import datetime -import json import logging -import os -from typing import Any -import aioboto3 # type: ignore[import-untyped] from database.connection import DatabaseRouter from database.models.control_plane import DatabaseShard from database.models.data_plane import DataPlaneOutbox @@ -14,6 +10,7 @@ from scheduler.ports.handler import JobHandlerPort from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession +from worker.ports.message_publisher import MessagePublisherPort logger = logging.getLogger(__name__) @@ -26,17 +23,14 @@ PipelineEventType.DELIVERY_COMPLETED: MessageQueueName.TRANSFORM_ORCHESTRATION_QUEUE, } -# Maximum number of events to sweep per run to bound wall-clock time _BATCH_SIZE = 100 _CONCURRENCY_LIMIT = 5 class DataPlaneOutboxSweeperJobHandler(JobHandlerPort): - def __init__(self, db_router: DatabaseRouter) -> None: + def __init__(self, db_router: DatabaseRouter, message_publisher: MessagePublisherPort) -> None: self.db_router = db_router - self._endpoint_url = os.environ.get("AWS_ENDPOINT_URL", "http://localhost:4566") - self._region = "us-east-1" - self._session = aioboto3.Session() + self.message_publisher = message_publisher async def execute(self, job: Job) -> datetime.datetime | None: """ @@ -48,45 +42,49 @@ async def execute(self, job: Job) -> datetime.datetime | None: total_processed = 0 sem = asyncio.Semaphore(_CONCURRENCY_LIMIT) - # We share one SQS client pool across all shards - async with self._session.client( - "sqs", endpoint_url=self._endpoint_url, region_name=self._region - ) as sqs: - queue_url_cache: dict[str, str] = {} - + async with self.message_publisher.connect(): async for global_session in self.db_router.get_global_session(): res = await global_session.execute(select(DatabaseShard)) shards = res.scalars().all() - async def _bounded_sweep(shard_name: str, shard_dsn: str) -> int: - async with sem: - return await self._sweep_shard(shard_name, shard_dsn, sqs, queue_url_cache) + async def _bounded_sweep(shard_name: str, shard_dsn: str) -> int: + async with sem: + try: + return await self._sweep_shard(shard_name, shard_dsn) + except Exception as e: + logger.error( + f"[DataPlaneOutboxSweeper] Failed sweeping shard {shard_name}: {e}" + ) + return 0 - results = await asyncio.gather( - *[_bounded_sweep(shard.name, shard.dsn) for shard in shards] - ) - total_processed += sum(results) + results = await asyncio.gather( + *[_bounded_sweep(shard.name, shard.dsn) for shard in shards] + ) + total_processed += sum(results) logger.info( f"[DataPlaneOutboxSweeper] Sweep complete. Total events forwarded: {total_processed}" ) - interval_seconds = job.payload.get("interval_seconds", 60) if job.payload else 60 + interval_seconds = ( + job.interval_seconds if job.interval_seconds and job.interval_seconds > 0 else 60 + ) return datetime.datetime.now(datetime.UTC) + datetime.timedelta(seconds=interval_seconds) - async def _sweep_shard( - self, shard_name: str, shard_dsn: str, sqs: Any, queue_url_cache: dict[str, str] - ) -> int: + async def _sweep_shard(self, shard_name: str, shard_dsn: str) -> int: """Sweep a single tenant shard outbox, dispatching via SQS batching.""" processed = 0 engine = await self.db_router.get_engine(shard_name, shard_dsn) async with AsyncSession(engine, expire_on_commit=False) as session: + # Only sweep events older than 5 minutes to avoid racing with Debezium CDC + five_mins_ago = datetime.datetime.now(datetime.UTC) - datetime.timedelta(minutes=5) stmt = ( select(DataPlaneOutbox) .where( DataPlaneOutbox.status == "PENDING", DataPlaneOutbox.event_type.in_(list(PipelineEventType)), + DataPlaneOutbox.created_at < five_mins_ago, ) .limit(_BATCH_SIZE) .with_for_update(skip_locked=True) @@ -112,58 +110,29 @@ async def _sweep_shard( batches_by_queue.setdefault(queue_name, []).append(event) for queue_name, queue_events in batches_by_queue.items(): - if queue_name not in queue_url_cache: - try: - resp = await sqs.get_queue_url(QueueName=queue_name) - queue_url_cache[queue_name] = resp["QueueUrl"] - except Exception: - logger.exception( - f"[DataPlaneOutboxSweeper] Failed to get queue url for {queue_name}" - ) - continue - - queue_url = queue_url_cache[queue_name] - - # SQS allows max 10 messages per batch - for i in range(0, len(queue_events), 10): - batch = queue_events[i : i + 10] - entries = [] - for event in batch: - entries.append( - { - "Id": str(event.id), # SQS entry ID must be string - "MessageBody": json.dumps( - { - "idempotency_key": str(event.idempotency_key), - "event_type": event.event_type, - "payload": event.payload, - "tenant_id": event.tenant_id, - } - ), - } - ) + messages = [] + for event in queue_events: + messages.append( + { + "Id": str(event.id), + "MessageBody": { + "idempotency_key": str(event.idempotency_key), + "event_type": event.event_type, + "payload": event.payload, + "tenant_id": event.tenant_id, + }, + } + ) - try: - resp = await sqs.send_message_batch(QueueUrl=queue_url, Entries=entries) - # Process successful IDs - for success in resp.get("Successful", []): - event_id = success["Id"] - # Find the event object - for ev in batch: - if str(ev.id) == event_id: - ev.status = "PROCESSED" - processed += 1 - break - - # Log failures if any - for failed in resp.get("Failed", []): - logger.error( - f"[DataPlaneOutboxSweeper] Failed to forward event id={failed['Id']}: " - f"{failed['Message']}" - ) - except Exception: - logger.exception( - f"[DataPlaneOutboxSweeper] Failed to send batch to {queue_name}" + successful_ids = await self.message_publisher.publish_batch(queue_name, messages) + + for event in queue_events: + if str(event.id) in successful_ids: + event.status = "PROCESSED" + processed += 1 + else: + logger.error( + f"[DataPlaneOutboxSweeper] Failed to forward event id={event.id} to {queue_name}" ) # Commit the session. Only events marked PROCESSED/FAILED above will be updated. diff --git a/services/workers/orchestrator/src/worker/ports/message_publisher.py b/services/workers/orchestrator/src/worker/ports/message_publisher.py new file mode 100644 index 00000000..7c357f3e --- /dev/null +++ b/services/workers/orchestrator/src/worker/ports/message_publisher.py @@ -0,0 +1,22 @@ +import abc +from contextlib import AbstractAsyncContextManager +from typing import Any + + +class MessagePublisherPort(abc.ABC): + @abc.abstractmethod + def connect(self) -> AbstractAsyncContextManager["MessagePublisherPort"]: + """ + Context manager to establish and share the underlying connection pool. + Must be entered before calling publish_batch. + """ + pass + + @abc.abstractmethod + async def publish_batch(self, queue_name: str, messages: list[dict[str, Any]]) -> list[str]: + """ + Publishes a batch of messages to the specified queue. + Each message dict MUST contain an "Id" key (string) used for batch deduplication. + Returns a list of the "Id"s that were successfully published. + """ + pass diff --git a/services/workers/orchestrator/tests/test_data_main.py b/services/workers/orchestrator/tests/test_data_main.py index 5838b133..33fc84d9 100644 --- a/services/workers/orchestrator/tests/test_data_main.py +++ b/services/workers/orchestrator/tests/test_data_main.py @@ -3,7 +3,11 @@ import pytest from database.connection import DatabaseRouter -from worker.data.main import poll_sqs_queue, process_pipeline_event, validate_target_url +from database.models.control_plane import DatabaseShard, Tenant +from worker.adapters.sqs_poller import poll_sqs_queue +from worker.core.security import validate_target_url +from worker.core.tenant_resolver import TenantResolver +from worker.data.handlers import process_pipeline_event GLOBAL_DB_URL = os.getenv( "DB_GLOBAL_URL", "postgresql+asyncpg://edi:edi_password@localhost:5432/edi_global" @@ -16,6 +20,13 @@ def test_validate_target_url(): assert validate_target_url("http://example.com") is True assert validate_target_url("http://127.0.0.1") is False + assert validate_target_url("ftp://example.com") is False + assert validate_target_url("http://") is False + assert validate_target_url("http://192.168.1.1") is False + assert validate_target_url("http://10.0.0.1") is False + + with patch("socket.getaddrinfo", side_effect=Exception("mock err")): + assert validate_target_url("http://example.com") is False @pytest.fixture @@ -49,8 +60,78 @@ async def test_process_pipeline_event_no_message(router: DatabaseRouter): @pytest.mark.asyncio @pytest.mark.integration +async def test_tenant_resolver_integration(router: DatabaseRouter): + """ + Real narrow integration test. Connects to the database and tests real behavior. + """ + global_gen = router.get_global_session() + global_session = await global_gen.__anext__() + + import uuid + + suffix = str(uuid.uuid4())[:8] + shard_name = f"test_shard_{suffix}" + tenant_name = f"Test Tenant_{suffix}" + + # 1. Insert a shard and a tenant into the real DB + shard = DatabaseShard( + name=shard_name, dsn="postgresql+asyncpg://test:pass@localhost:5433/test_shard" + ) + global_session.add(shard) + await global_session.commit() + + tenant = Tenant( + name=tenant_name, shard_id=shard.id, tier="standard", shard_schema=f"tenant_test_{suffix}" + ) + global_session.add(tenant) + await global_session.commit() + + tenant_id = tenant.id + await global_gen.aclose() + + # 2. Use the resolver + resolver = TenantResolver(db_router=router, ttl_secs=300) + + try: + # Resolving once should hit DB + resolved_shard_name, shard_dsn = await resolver.resolve(tenant_id) + assert resolved_shard_name == shard_name + assert shard_dsn == "postgresql+asyncpg://test:pass@localhost:5433/test_shard" + + # Resolving again should hit cache + resolved_shard_name2, shard_dsn2 = await resolver.resolve(tenant_id) + assert resolved_shard_name2 == shard_name + finally: + # Cleanup + global_gen2 = router.get_global_session() + global_session2 = await global_gen2.__anext__() + + tenant_to_delete = await global_session2.get(Tenant, tenant_id) + if tenant_to_delete: + await global_session2.delete(tenant_to_delete) + await global_session2.flush() + + shard_to_delete = await global_session2.get(DatabaseShard, shard.id) + if shard_to_delete: + await global_session2.delete(shard_to_delete) + + await global_session2.commit() + await global_gen2.aclose() + + +@pytest.mark.asyncio +@pytest.mark.integration +async def test_tenant_resolver_not_found(router: DatabaseRouter): + """ + Test TenantResolver when a tenant is not found in the live DB. + """ + resolver = TenantResolver(db_router=router, ttl_secs=300) + with pytest.raises(ValueError, match="Tenant -999 not found in Global DB"): + await resolver.resolve(-999) + + async def test_process_delivery_no_message(router: DatabaseRouter): - from worker.data.main import process_delivery + from worker.data.handlers import process_delivery resolver = AsyncMock() resolver.resolve.return_value = ("shard_1", SHARD_1_URL) @@ -100,25 +181,25 @@ async def __aexit__(self, exc_type, exc_val, exc_tb): mock_processor = AsyncMock() with ( - patch("worker.data.main.aioboto3.Session", return_value=mock_session), - patch("worker.data.main.asyncio.sleep", side_effect=Exception("Break out of retry loop")), + patch("worker.adapters.sqs_poller.aioboto3.Session", return_value=mock_session), + patch( + "worker.adapters.sqs_poller.asyncio.sleep", + side_effect=Exception("Break out of retry loop"), + ), ): try: await poll_sqs_queue( "test-queue", processor_func=mock_processor, - resolver=AsyncMock(), - db_router=AsyncMock(), - s3_bucket="test-bucket", aws_endpoint=None, ) except Exception as e: if str(e) != "Break out of retry loop": raise - # Ensure processor was called for the valid message - mock_processor.assert_called_once() - assert mock_processor.call_args[1]["trace_id"] == "123" + # Ensure processor was called for both valid and poison message (since validation moved to main.py wrappers) + assert mock_processor.call_count == 2 + assert mock_processor.call_args_list[0][0][0]["payload"]["trace_id"] == "123" # Ensure all 3 messages were deleted (1 success, 2 poison) assert mock_sqs.delete_message.call_count == 3 diff --git a/services/workers/orchestrator/tests/test_sqs_publisher.py b/services/workers/orchestrator/tests/test_sqs_publisher.py new file mode 100644 index 00000000..22af4720 --- /dev/null +++ b/services/workers/orchestrator/tests/test_sqs_publisher.py @@ -0,0 +1,93 @@ +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from worker.adapters.sqs_publisher import SqsPublisherAdapter + +pytestmark = pytest.mark.asyncio + + +async def test_sqs_publisher_publish_batch(): + mock_session = MagicMock() + mock_client = AsyncMock() + mock_client_ctx = MagicMock() + mock_client_ctx.__aenter__ = AsyncMock(return_value=mock_client) + mock_client_ctx.__aexit__ = AsyncMock() + mock_session.client.return_value = mock_client_ctx + + # Mock get_queue_url + mock_client.get_queue_url.return_value = {"QueueUrl": "http://sqs/test"} + # Mock send_message_batch + mock_client.send_message_batch.return_value = { + "Successful": [{"Id": "1"}, {"Id": "2"}], + "Failed": [{"Id": "3", "Message": "Error"}], + } + + with patch("aioboto3.Session", return_value=mock_session): + adapter = SqsPublisherAdapter(region="us-east-1", endpoint_url=None) + + async with adapter.connect(): + # Publish messages + messages = [ + {"Id": "1", "MessageBody": {"event": "A"}}, + {"Id": "2", "MessageBody": {"event": "B"}}, + {"Id": "3", "MessageBody": {"event": "C"}}, + {"MessageBody": {"event": "MissingId"}}, # Should be skipped + ] + + success_ids = await adapter.publish_batch("test-queue", messages) + + assert success_ids == ["1", "2"] + mock_client.get_queue_url.assert_called_once_with(QueueName="test-queue") + mock_client.send_message_batch.assert_called_once() + + +async def test_sqs_publisher_not_connected(): + adapter = SqsPublisherAdapter(region="us-east-1", endpoint_url=None) + with pytest.raises(RuntimeError, match="must be called within the connect"): + await adapter.publish_batch("test-queue", [{"Id": "1"}]) + + +async def test_sqs_publisher_get_queue_url_error(): + mock_session = MagicMock() + mock_client = AsyncMock() + mock_client_ctx = MagicMock() + mock_client_ctx.__aenter__ = AsyncMock(return_value=mock_client) + mock_client_ctx.__aexit__ = AsyncMock() + mock_session.client.return_value = mock_client_ctx + + mock_client.get_queue_url.side_effect = Exception("SQS Error") + + with patch("aioboto3.Session", return_value=mock_session): + adapter = SqsPublisherAdapter(region="us-east-1", endpoint_url=None) + + async with adapter.connect(): + success_ids = await adapter.publish_batch( + "test-queue", [{"Id": "1", "MessageBody": {}}] + ) + assert success_ids == [] + + +async def test_sqs_publisher_publish(): + mock_session = MagicMock() + mock_client = AsyncMock() + mock_client_ctx = MagicMock() + mock_client_ctx.__aenter__ = AsyncMock(return_value=mock_client) + mock_client_ctx.__aexit__ = AsyncMock() + mock_session.client.return_value = mock_client_ctx + + mock_client.get_queue_url.return_value = {"QueueUrl": "http://sqs/test"} + + with patch("aioboto3.Session", return_value=mock_session): + adapter = SqsPublisherAdapter(region="us-east-1", endpoint_url=None) + await adapter.publish("test-queue", {"event": "A"}) + + # Verify get_queue_url and send_message were called + mock_client.get_queue_url.assert_called_once_with(QueueName="test-queue") + mock_client.send_message.assert_called_once_with( + QueueUrl="http://sqs/test", MessageBody='{"event": "A"}' + ) + + # Call again to verify cache is used (get_queue_url should not be called again) + await adapter.publish("test-queue", {"event": "B"}) + assert mock_client.get_queue_url.call_count == 1 + assert mock_client.send_message.call_count == 2 diff --git a/uv.lock b/uv.lock index dd7ccc09..7ea3496c 100644 --- a/uv.lock +++ b/uv.lock @@ -897,6 +897,18 @@ toml = [ { name = "tomli", marker = "python_full_version <= '3.11'" }, ] +[[package]] +name = "croniter" +version = "6.2.4" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "python-dateutil" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/37/57/2e2a65aee2a70483cb28e2b7e15a072d00a523207593b44400d4717bb100/croniter-6.2.4.tar.gz", hash = "sha256:fc124f751b1b04805c2a04b061898b436b45ab2320b045e1e052ea895de65189", size = 166267, upload-time = "2026-07-10T09:52:59.955Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/cd/ba/d678e5bd329646ca51d3c92addbc77804e86d21f4b6b6a027218e6abb010/croniter-6.2.4-py3-none-any.whl", hash = "sha256:8ef3d544107a5c05a150a2d78f8bf5a8eb9c5c4d93405a736b824109574e3f4d", size = 46677, upload-time = "2026-07-10T09:52:58.425Z" }, +] + [[package]] name = "cryptography" version = "49.0.0" @@ -2922,6 +2934,7 @@ name = "scheduler" version = "0.1.0" source = { editable = "libs/scheduler" } dependencies = [ + { name = "croniter" }, { name = "database" }, { name = "pydantic" }, { name = "sqlalchemy" }, @@ -2929,6 +2942,7 @@ dependencies = [ [package.metadata] requires-dist = [ + { name = "croniter", specifier = ">=6.2.4" }, { name = "database", editable = "libs/database" }, { name = "pydantic", specifier = ">=2.0.0" }, { name = "sqlalchemy", specifier = ">=2.0.0" }, From 58b514563e2cad2249d14e0241ea5e71c344b250 Mon Sep 17 00:00:00 2001 From: Pramod Date: Fri, 17 Jul 2026 21:46:28 +0530 Subject: [PATCH 4/5] light weight scheduler and outbox sweeper code review --- .../components/CreatePartnerModal.tsx | 71 ++++--- .../components/SchedulerDashboard.tsx | 4 +- libs/pipeline/src/pipeline/adapters/as2.py | 27 ++- libs/pipeline/src/pipeline/adapters/http.py | 31 ++- .../src/pipeline/core/delivery/as2.py | 8 +- .../src/pipeline/core/delivery/base.py | 8 +- .../src/pipeline/core/delivery/router.py | 4 +- .../src/pipeline/core/delivery/sftp.py | 8 +- .../src/pipeline/core/delivery/webhook.py | 14 +- libs/pipeline/src/pipeline/ports/http.py | 6 +- libs/pipeline/tests/fakes.py | 15 +- .../src/scheduler/adapters/repository.py | 33 ++- libs/scheduler/src/scheduler/core/service.py | 10 + .../src/scheduler/ports/repository.py | 6 +- pyproject.toml | 4 +- .../api/src/api/routers/platform/scheduler.py | 20 +- .../trading_partners/platform/as2_partners.py | 9 +- .../api/tests/test_api_receiver_service.py | 39 ++++ services/api/tests/test_scheduler.py | 6 +- services/as2_server/scripts/seed.py | 2 - .../src/worker/adapters/sqs_poller.py | 9 +- .../src/worker/adapters/sqs_publisher.py | 3 +- .../orchestrator/src/worker/core/security.py | 54 +++++ .../src/worker/core/tenant_resolver.py | 17 +- .../orchestrator/src/worker/data/handlers.py | 53 ++++- .../orchestrator/src/worker/data/main.py | 24 ++- .../src/worker/data/scheduled_jobs_handler.py | 2 +- .../src/worker/jobs/data_retention.py | 2 +- .../workers/orchestrator/src/worker/main.py | 11 +- .../orchestrator/tests/test_data_main.py | 30 ++- .../orchestrator/tests/test_data_retention.py | 91 ++++++++ .../orchestrator/tests/test_handlers.py | 195 ++++++++++++++++++ .../orchestrator/tests/test_outbox_sweeper.py | 141 +++++++++++++ .../tests/test_scheduled_jobs_handler.py | 51 +++++ .../orchestrator/tests/test_security.py | 58 ++++++ 35 files changed, 969 insertions(+), 97 deletions(-) create mode 100644 services/workers/orchestrator/tests/test_data_retention.py create mode 100644 services/workers/orchestrator/tests/test_handlers.py create mode 100644 services/workers/orchestrator/tests/test_outbox_sweeper.py create mode 100644 services/workers/orchestrator/tests/test_scheduled_jobs_handler.py create mode 100644 services/workers/orchestrator/tests/test_security.py diff --git a/frontend/web/src/features/partners/components/CreatePartnerModal.tsx b/frontend/web/src/features/partners/components/CreatePartnerModal.tsx index 3ce2974e..7f05216d 100644 --- a/frontend/web/src/features/partners/components/CreatePartnerModal.tsx +++ b/frontend/web/src/features/partners/components/CreatePartnerModal.tsx @@ -30,24 +30,23 @@ export function CreatePartnerModal({ existingAs2Ids = [] }: { existingAs2Ids?: s const handleCleanup = async () => { if (privateKeyVaultRef) { - try { - await deleteCertSecret.mutateAsync(privateKeyVaultRef); - } catch (e) { - console.error("Failed to cleanup orphaned secret", e); - } + await deleteCertSecret.mutateAsync(privateKeyVaultRef); } }; - const reset = () => { + const reset = async () => { // Only cleanup if we are abandoning an unsaved draft - handleCleanup().then(() => { - setIsLocal(false); - setCertPem(''); + try { + await handleCleanup(); setPrivateKeyVaultRef(null); + setCertPem(''); setGeneratedForAs2Id(null); + setIsLocal(false); setAs2Id(''); setUrl(''); - }); + } catch (e) { + console.error("Failed to cleanup orphaned secret during reset", e); + } }; const handleOpenChange = (open: boolean) => { @@ -87,12 +86,17 @@ export function CreatePartnerModal({ existingAs2Ids = [] }: { existingAs2Ids?: s if (isLocal && privateKeyVaultRef && generatedForAs2Id && submittedAs2Id !== generatedForAs2Id) { // Invalidate existing if AS2 ID changed - await handleCleanup(); + try { + await handleCleanup(); + setPrivateKeyVaultRef(null); + setCertPem(''); + setGeneratedForAs2Id(null); + } catch (e) { + toast({ title: 'Error', description: 'Failed to cleanup old certificate.', variant: 'destructive' }); + return; + } finalCertPem = ''; finalVaultRef = null; - setCertPem(''); - setPrivateKeyVaultRef(null); - setGeneratedForAs2Id(null); toast({ title: 'Warning', description: 'AS2 ID changed. Please regenerate the certificate.', variant: 'destructive' }); return; } @@ -147,20 +151,31 @@ export function CreatePartnerModal({ existingAs2Ids = [] }: { existingAs2Ids?: s aria-checked={isLocal} onClick={async () => { const nextIsLocal = !isLocal; - setIsLocal(nextIsLocal); if (nextIsLocal) { - await handleCleanup(); - setCertPem(''); - setPrivateKeyVaultRef(null); - setGeneratedForAs2Id(null); + try { + await handleCleanup(); + setPrivateKeyVaultRef(null); + setCertPem(''); + setGeneratedForAs2Id(null); + } catch (e) { + toast({ title: 'Error', description: 'Failed to cleanup old certificate.', variant: 'destructive' }); + return; + } + setIsLocal(nextIsLocal); if (!url && platformSettings?.available_as2_receive_urls?.length) { setUrl(platformSettings.available_as2_receive_urls[0]); } } else { - await handleCleanup(); - setCertPem(''); - setPrivateKeyVaultRef(null); - setGeneratedForAs2Id(null); + try { + await handleCleanup(); + setPrivateKeyVaultRef(null); + setCertPem(''); + setGeneratedForAs2Id(null); + } catch (e) { + toast({ title: 'Error', description: 'Failed to cleanup old certificate.', variant: 'destructive' }); + return; + } + setIsLocal(nextIsLocal); if (platformSettings?.available_as2_receive_urls?.includes(url)) { setUrl(''); } @@ -245,7 +260,15 @@ export function CreatePartnerModal({ existingAs2Ids = [] }: { existingAs2Ids?: s return; } if (privateKeyVaultRef) { - await handleCleanup(); + try { + await handleCleanup(); + setPrivateKeyVaultRef(null); + setCertPem(''); + setGeneratedForAs2Id(null); + } catch (e) { + toast({ title: 'Error', description: 'Failed to cleanup old certificate.', variant: 'destructive' }); + return; + } } generateCert.mutate(as2Id, { onSuccess: (res) => { diff --git a/frontend/web/src/features/platform/components/SchedulerDashboard.tsx b/frontend/web/src/features/platform/components/SchedulerDashboard.tsx index e710270d..ef030d8e 100644 --- a/frontend/web/src/features/platform/components/SchedulerDashboard.tsx +++ b/frontend/web/src/features/platform/components/SchedulerDashboard.tsx @@ -69,11 +69,11 @@ export const SchedulerDashboard = () => { try { if (scheduleType === 'interval') { - const val = parseInt(intervalValue, 10); - if (isNaN(val) || val <= 0) { + if (!/^\d+$/.test(intervalValue) || parseInt(intervalValue, 10) <= 0) { setIntervalError('Interval must be a positive integer.'); return; } + const val = parseInt(intervalValue, 10); let multiplier = 1; if (intervalUnit === 'minutes') multiplier = 60; diff --git a/libs/pipeline/src/pipeline/adapters/as2.py b/libs/pipeline/src/pipeline/adapters/as2.py index 2113b9bc..e55cc120 100644 --- a/libs/pipeline/src/pipeline/adapters/as2.py +++ b/libs/pipeline/src/pipeline/adapters/as2.py @@ -4,6 +4,7 @@ """ import logging +import typing import httpx from pipeline.ports.as2 import AS2DeliveryPort @@ -20,8 +21,11 @@ class HttpxAS2DeliveryAdapter(AS2DeliveryPort): connection errors. MDN parsing is delegated to the caller (DeliveryService). """ - def __init__(self, timeout_secs: int = 30) -> None: + def __init__( + self, timeout_secs: int = 30, validator: typing.Callable[[str], typing.Any] | None = None + ) -> None: self.timeout = timeout_secs + self.validator = validator async def deliver( self, @@ -38,11 +42,22 @@ async def deliver( """ logger.debug(f"AS2 HTTP POST → {url}, Content-Length={len(body)}") - async with httpx.AsyncClient( - timeout=self.timeout, - follow_redirects=False, # AS2 spec does not permit redirect following - ) as client: - response = await client.post(url, content=body, headers=headers) + import contextlib + + ctx = self.validator(url) if self.validator else contextlib.nullcontext() + + # If validator returns a boolean (legacy), handle it + if isinstance(ctx, bool): + if not ctx: + raise ValueError("URL validation failed for provided destination.") + ctx = contextlib.nullcontext() + + with ctx: + async with httpx.AsyncClient( + timeout=self.timeout, + follow_redirects=False, # AS2 spec does not permit redirect following + ) as client: + response = await client.post(url, content=body, headers=headers) # Convert httpx headers to a standard dict resp_headers = {k.lower(): v for k, v in response.headers.items()} diff --git a/libs/pipeline/src/pipeline/adapters/http.py b/libs/pipeline/src/pipeline/adapters/http.py index 7a50c1f8..85324d81 100644 --- a/libs/pipeline/src/pipeline/adapters/http.py +++ b/libs/pipeline/src/pipeline/adapters/http.py @@ -1,4 +1,5 @@ from collections.abc import Callable +from typing import Any import httpx from pipeline.ports.http import HttpDeliveryPort @@ -9,20 +10,34 @@ class HttpxDeliveryAdapter(HttpDeliveryPort): Concrete implementation of HttpDeliveryPort using HTTPX. """ - def __init__(self, timeout_secs: int = 30, validator: Callable[[str], bool] | None = None): + def __init__(self, timeout_secs: int = 30, validator: Callable[[str], Any] | None = None): self.timeout = timeout_secs self.validator = validator async def deliver( - self, url: str, payload: bytes, auth_token: str | None = None + self, + url: str, + payload: bytes, + auth_token: str | None = None, + idempotency_key: str | None = None, ) -> tuple[int, str]: - if self.validator and not self.validator(url): - raise ValueError("URL validation failed for provided destination.") - headers = {"Content-Type": "application/json"} if auth_token: headers["Authorization"] = auth_token + if idempotency_key: + headers["Idempotency-Key"] = idempotency_key + + import contextlib + + ctx = self.validator(url) if self.validator else contextlib.nullcontext() + + # If validator returns a boolean (legacy), handle it + if isinstance(ctx, bool): + if not ctx: + raise ValueError("URL validation failed for provided destination.") + ctx = contextlib.nullcontext() - async with httpx.AsyncClient(timeout=self.timeout, follow_redirects=False) as client: - response = await client.post(url, content=payload, headers=headers) - return response.status_code, response.text + with ctx: + async with httpx.AsyncClient(timeout=self.timeout, follow_redirects=False) as client: + response = await client.post(url, content=payload, headers=headers) + return response.status_code, response.text diff --git a/libs/pipeline/src/pipeline/core/delivery/as2.py b/libs/pipeline/src/pipeline/core/delivery/as2.py index 3b231af4..0c720169 100644 --- a/libs/pipeline/src/pipeline/core/delivery/as2.py +++ b/libs/pipeline/src/pipeline/core/delivery/as2.py @@ -22,7 +22,13 @@ def __init__( self.as2_delivery = as2_delivery self._as2_orchestrator = AS2MessageOrchestrator(vault=vault) - async def deliver(self, trace_id: str, partner_id: str, edi_msg: EdiMessageDomainModel) -> None: + async def deliver( + self, + trace_id: str, + partner_id: str, + edi_msg: EdiMessageDomainModel, + idempotency_key: str | None = None, + ) -> None: if not await self.repository.claim_edi_message(trace_id): logger.warning(f"Could not claim trace_id={trace_id} (already claimed or terminal).") return diff --git a/libs/pipeline/src/pipeline/core/delivery/base.py b/libs/pipeline/src/pipeline/core/delivery/base.py index b6bb68ea..d60fe36f 100644 --- a/libs/pipeline/src/pipeline/core/delivery/base.py +++ b/libs/pipeline/src/pipeline/core/delivery/base.py @@ -28,5 +28,11 @@ async def _emit_delivery_completed(self, trace_id: str, direction: str, status: }, ) - async def deliver(self, trace_id: str, partner_id: str, edi_msg: EdiMessageDomainModel) -> None: + async def deliver( + self, + trace_id: str, + partner_id: str, + edi_msg: EdiMessageDomainModel, + idempotency_key: str | None = None, + ) -> None: raise NotImplementedError diff --git a/libs/pipeline/src/pipeline/core/delivery/router.py b/libs/pipeline/src/pipeline/core/delivery/router.py index 0377fb95..93c1037f 100644 --- a/libs/pipeline/src/pipeline/core/delivery/router.py +++ b/libs/pipeline/src/pipeline/core/delivery/router.py @@ -19,7 +19,7 @@ def __init__( self.repository = repository self.strategies = strategies - async def deliver(self, trace_id: str) -> None: + async def deliver(self, trace_id: str, idempotency_key: str | None = None) -> None: """ Looks up the route for the given trace_id and dispatches to the correct delivery handler via the strategy registry. @@ -65,7 +65,7 @@ async def deliver(self, trace_id: str) -> None: for route_key, strategy in self.strategies.items(): partner_id = route.get(route_key) if partner_id: - await strategy.deliver(trace_id, partner_id, edi_msg) + await strategy.deliver(trace_id, partner_id, edi_msg, idempotency_key) return raise ValueError( diff --git a/libs/pipeline/src/pipeline/core/delivery/sftp.py b/libs/pipeline/src/pipeline/core/delivery/sftp.py index 6f6cabfb..51d9b404 100644 --- a/libs/pipeline/src/pipeline/core/delivery/sftp.py +++ b/libs/pipeline/src/pipeline/core/delivery/sftp.py @@ -20,7 +20,13 @@ def __init__( super().__init__(repository, vault) self.sftp_delivery = sftp_delivery - async def deliver(self, trace_id: str, partner_id: str, edi_msg: EdiMessageDomainModel) -> None: + async def deliver( + self, + trace_id: str, + partner_id: str, + edi_msg: EdiMessageDomainModel, + idempotency_key: str | None = None, + ) -> None: if not await self.repository.claim_edi_message(trace_id): logger.warning(f"Could not claim trace_id={trace_id} (already claimed or terminal).") return diff --git a/libs/pipeline/src/pipeline/core/delivery/webhook.py b/libs/pipeline/src/pipeline/core/delivery/webhook.py index 30646b23..861b21f8 100644 --- a/libs/pipeline/src/pipeline/core/delivery/webhook.py +++ b/libs/pipeline/src/pipeline/core/delivery/webhook.py @@ -21,7 +21,13 @@ def __init__( super().__init__(repository, vault) self.http_delivery = http_delivery - async def deliver(self, trace_id: str, partner_id: str, edi_msg: EdiMessageDomainModel) -> None: + async def deliver( + self, + trace_id: str, + partner_id: str, + edi_msg: EdiMessageDomainModel, + idempotency_key: str | None = None, + ) -> None: if not await self.repository.claim_api_payload(trace_id): logger.warning(f"Could not claim trace_id={trace_id} (already claimed or terminal).") return @@ -49,8 +55,12 @@ async def deliver(self, trace_id: str, partner_id: str, edi_msg: EdiMessageDomai ) auth_token = await self.vault.get_secret(partner["auth_header_vault_ref"]) + # Pass idempotency_key down to the http_delivery if it supports it, or add it to headers manually status_code, response_text = await self.http_delivery.deliver( - url=partner["url"], payload=raw_payload, auth_token=auth_token + url=partner["url"], + payload=raw_payload, + auth_token=auth_token, + idempotency_key=idempotency_key, ) except Exception as e: await self.repository.update_api_payload_status( diff --git a/libs/pipeline/src/pipeline/ports/http.py b/libs/pipeline/src/pipeline/ports/http.py index 82efe946..8a691f20 100644 --- a/libs/pipeline/src/pipeline/ports/http.py +++ b/libs/pipeline/src/pipeline/ports/http.py @@ -7,7 +7,11 @@ class HttpDeliveryPort(Protocol): """ async def deliver( - self, url: str, payload: bytes, auth_token: str | None = None + self, + url: str, + payload: bytes, + auth_token: str | None = None, + idempotency_key: str | None = None, ) -> tuple[int, str]: """Sends the payload to the specified URL and returns the HTTP status code and response body.""" ... diff --git a/libs/pipeline/tests/fakes.py b/libs/pipeline/tests/fakes.py index e299c21f..59731df4 100644 --- a/libs/pipeline/tests/fakes.py +++ b/libs/pipeline/tests/fakes.py @@ -233,9 +233,20 @@ def __init__(self, status_code: int = 200) -> None: self.status_code = status_code async def deliver( - self, url: str, payload: bytes, auth_token: str | None = None + self, + url: str, + payload: bytes, + auth_token: str | None = None, + idempotency_key: str | None = None, ) -> tuple[int, str]: - self.delivered.append({"url": url, "payload": payload, "auth_token": auth_token}) + self.delivered.append( + { + "url": url, + "payload": payload, + "auth_token": auth_token, + "idempotency_key": idempotency_key, + } + ) return self.status_code, "Mock response body" diff --git a/libs/scheduler/src/scheduler/adapters/repository.py b/libs/scheduler/src/scheduler/adapters/repository.py index e7f151fa..e985b860 100644 --- a/libs/scheduler/src/scheduler/adapters/repository.py +++ b/libs/scheduler/src/scheduler/adapters/repository.py @@ -13,8 +13,9 @@ class SqlAlchemyJobRepository(JobRepositoryPort): - def __init__(self, engine: AsyncEngine): + def __init__(self, engine: AsyncEngine, lock_lease_seconds: int = 300): self.session_factory = async_sessionmaker(engine, expire_on_commit=False) + self.lock_lease_seconds = lock_lease_seconds def _to_domain(self, record: ScheduledJob) -> Job: return Job( @@ -47,7 +48,16 @@ async def claim_next_jobs(self, worker_id: str, limit: int) -> list[Job]: stmt = ( select(ScheduledJob) .where( - (ScheduledJob.status == JobStatus.PENDING.value) + ( + (ScheduledJob.status == JobStatus.PENDING.value) + | ( + (ScheduledJob.status == JobStatus.RUNNING.value) + & ( + ScheduledJob.locked_at + < now - datetime.timedelta(seconds=self.lock_lease_seconds) + ) + ) + ) & (ScheduledJob.next_run_at.is_(None) | (ScheduledJob.next_run_at <= now)) ) .order_by( @@ -73,6 +83,25 @@ async def claim_next_jobs(self, worker_id: str, limit: int) -> list[Job]: await session.flush() return claimed + async def sweep_stuck_jobs(self, timeout: datetime.timedelta) -> int: + now = datetime.datetime.now(datetime.UTC) + threshold = now - timeout + async with self.session_factory() as session, session.begin(): + stmt = ( + update(ScheduledJob) + .where( + (ScheduledJob.status == JobStatus.RUNNING.value) + & (ScheduledJob.locked_at < threshold) + ) + .values( + status=JobStatus.PENDING.value, + locked_at=None, + locked_by=None, + ) + ) + result = await session.execute(stmt) + return int(result.rowcount) if result.rowcount is not None else 0 # type: ignore[attr-defined] + async def mark_completed(self, job_id: uuid.UUID) -> None: async with self.session_factory() as session, session.begin(): stmt = ( diff --git a/libs/scheduler/src/scheduler/core/service.py b/libs/scheduler/src/scheduler/core/service.py index d3f533ab..3a3ce1c7 100644 --- a/libs/scheduler/src/scheduler/core/service.py +++ b/libs/scheduler/src/scheduler/core/service.py @@ -89,8 +89,18 @@ async def _execute_job(self, job: Any) -> None: await self.repository.mark_failed(job.id, error=str(e)) async def _poll_loop(self, poll_interval_seconds: float) -> None: + import datetime + + last_sweep = datetime.datetime.now(datetime.UTC) while self._is_running: try: + now = datetime.datetime.now(datetime.UTC) + if (now - last_sweep).total_seconds() > 60: + swept = await self.repository.sweep_stuck_jobs(datetime.timedelta(seconds=300)) + if swept > 0: + logger.info(f"Swept {swept} stuck jobs back to PENDING.") + last_sweep = now + # Remove completed tasks from active set self._active_jobs = {task for task in self._active_jobs if not task.done()} diff --git a/libs/scheduler/src/scheduler/ports/repository.py b/libs/scheduler/src/scheduler/ports/repository.py index 16c88a11..b69ea9a4 100644 --- a/libs/scheduler/src/scheduler/ports/repository.py +++ b/libs/scheduler/src/scheduler/ports/repository.py @@ -1,6 +1,6 @@ import abc import uuid -from datetime import datetime +from datetime import datetime, timedelta from typing import Any from scheduler.domain.models import Job @@ -11,6 +11,10 @@ class JobRepositoryPort(abc.ABC): async def claim_next_jobs(self, worker_id: str, limit: int) -> list[Job]: pass + @abc.abstractmethod + async def sweep_stuck_jobs(self, timeout: timedelta) -> int: + pass + @abc.abstractmethod async def mark_completed(self, job_id: uuid.UUID) -> None: pass diff --git a/pyproject.toml b/pyproject.toml index faad3f58..eed93904 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -127,6 +127,6 @@ omit = [ ] [tool.coverage.report] -# Required 75% coverage +# Required 80% coverage show_missing = true -fail_under = 75 +fail_under = 80 diff --git a/services/api/src/api/routers/platform/scheduler.py b/services/api/src/api/routers/platform/scheduler.py index 4ae81c5e..0ab70487 100644 --- a/services/api/src/api/routers/platform/scheduler.py +++ b/services/api/src/api/routers/platform/scheduler.py @@ -77,6 +77,20 @@ async def create_job(request: JobCreateRequest, uow: UnitOfWork = Depends(get_uo from sqlalchemy.exc import IntegrityError now = datetime.now(UTC) + + from scheduler.registry import SYSTEM_JOB_REGISTRY + + job_def = next((j for j in SYSTEM_JOB_REGISTRY if j.name.value == request.name), None) + if not job_def: + raise HTTPException( + status_code=422, detail=f"Job '{request.name}' is not a registered system job." + ) + + if not job_def.target_queue: + raise HTTPException( + status_code=422, detail=f"Job '{request.name}' has no configured target queue." + ) + async with uow: try: async with uow.global_session.begin_nested(): @@ -87,8 +101,12 @@ async def create_job(request: JobCreateRequest, uow: UnitOfWork = Depends(get_uo interval_seconds=request.interval_seconds, status=JobStatus.PENDING.value, next_run_at=now, + target_queue=job_def.target_queue, + app_namespace=job_def.app_namespace, + min_interval_seconds=job_def.min_interval_seconds, + max_interval_seconds=job_def.max_interval_seconds, retry_count=0, - max_retries=3, + max_retries=job_def.max_retries, created_at=now, updated_at=now, ) 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 9f667455..ee9a3b7a 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 @@ -96,13 +96,12 @@ async def create_platform_as2_partner( Creates a new Global AS2 Trading Partner (Local or Remote) in the Control Plane. If is_local is True, automatically generates a self-signed cert and stores private key in Vault. """ + public_cert_pem = request.public_cert_pem + private_key_vault_ref = request.private_key_vault_ref + auto_generated = False + try: async with uow: - public_cert_pem = request.public_cert_pem - private_key_vault_ref = request.private_key_vault_ref - - auto_generated = False - if request.is_local: if private_key_vault_ref: # Pre-stored vault ref (from generate cert flow) — use as-is diff --git a/services/api/tests/test_api_receiver_service.py b/services/api/tests/test_api_receiver_service.py index a7876602..a2c463d7 100644 --- a/services/api/tests/test_api_receiver_service.py +++ b/services/api/tests/test_api_receiver_service.py @@ -24,3 +24,42 @@ async def test_process_api_edi_json_success(): assert kwargs["event_type"] == PipelineEventType.TRANSFORM_EVENT assert kwargs["idempotency_key"] == trace_id + + +@pytest.mark.asyncio +async def test_process_api_edi_json_heading(): + mock_uow = AsyncMock() + svc = ApiReceiverService(mock_uow) + trace_id = await svc.process_api_edi_json( + tenant_id=1, + trading_partner_id="PARTNER_X", + payload=[ + {"heading": {"transaction_set_header_ST": {"transaction_set_identifier_code": "850"}}} + ], + ) + assert trace_id is not None + mock_uow.transactions.create_edi_json.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_process_api_edi_json_st_segment(): + mock_uow = AsyncMock() + svc = ApiReceiverService(mock_uow) + trace_id = await svc.process_api_edi_json( + tenant_id=1, trading_partner_id="PARTNER_X", payload=[{"ST": {"ST01": "855"}}] + ) + assert trace_id is not None + mock_uow.transactions.create_edi_json.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_process_api_edi_json_list_extraction(): + mock_uow = AsyncMock() + svc = ApiReceiverService(mock_uow) + # Give a list payload to hit the extraction logic for lists + payload = [{"transaction_type": "850", "foo": "bar"}, {"transaction_type": "850", "foo": "baz"}] + trace_id = await svc.process_api_edi_json( + tenant_id=1, trading_partner_id="PARTNER_X", payload=payload + ) + assert trace_id is not None + mock_uow.transactions.create_edi_json.assert_awaited_once() diff --git a/services/api/tests/test_scheduler.py b/services/api/tests/test_scheduler.py index 7ae24b82..7fec8d4e 100644 --- a/services/api/tests/test_scheduler.py +++ b/services/api/tests/test_scheduler.py @@ -91,11 +91,11 @@ def test_scheduler_update_job_invalid_interval(client, mock_uow): def test_scheduler_create_job(client, mock_uow): resp = client.post( "/api/v1/platform/scheduler/jobs", - json={"name": "test_job", "interval_seconds": 60, "payload": {}}, + json={"name": "outbox_sweeper", "interval_seconds": 60, "payload": {}}, ) assert resp.status_code == 200 assert mock_uow.global_session.add.called - assert resp.json()["name"] == "test_job" + assert resp.json()["name"] == "outbox_sweeper" def test_scheduler_create_job_already_exists(client, mock_uow): @@ -105,7 +105,7 @@ def test_scheduler_create_job_already_exists(client, mock_uow): resp = client.post( "/api/v1/platform/scheduler/jobs", - json={"name": "test_job"}, + json={"name": "outbox_sweeper"}, ) assert resp.status_code == 409 diff --git a/services/as2_server/scripts/seed.py b/services/as2_server/scripts/seed.py index 5642a854..e2b32c23 100644 --- a/services/as2_server/scripts/seed.py +++ b/services/as2_server/scripts/seed.py @@ -135,8 +135,6 @@ async def seed_database() -> None: # ensuring existing rows stay consistent after schema migrations. job.target_queue = job_def.target_queue job.app_namespace = job_def.app_namespace - job.cron_expression = job_def.default_cron_expression - job.timezone = job_def.default_timezone job.min_interval_seconds = job_def.min_interval_seconds job.max_interval_seconds = job_def.max_interval_seconds logger.info(f"Synced config for system job: {job_def.name}.") diff --git a/services/workers/orchestrator/src/worker/adapters/sqs_poller.py b/services/workers/orchestrator/src/worker/adapters/sqs_poller.py index a7c4d781..dac0c6cb 100644 --- a/services/workers/orchestrator/src/worker/adapters/sqs_poller.py +++ b/services/workers/orchestrator/src/worker/adapters/sqs_poller.py @@ -70,14 +70,7 @@ async def poll_sqs_queue( await sqs.delete_message( QueueUrl=queue_url, ReceiptHandle=receipt_handle ) - except (KeyError, NotImplementedError) as e: - # Permanently delete messages with unrecoverable configuration errors - logger.error( - f"[{queue_name}] Permanent validation error, deleting permanently: {e}" - ) - await sqs.delete_message( - QueueUrl=queue_url, ReceiptHandle=receipt_handle - ) + except Exception as e: logger.exception( f"[{queue_name}] Transient error processing message: {e}" diff --git a/services/workers/orchestrator/src/worker/adapters/sqs_publisher.py b/services/workers/orchestrator/src/worker/adapters/sqs_publisher.py index a741c13b..32e53091 100644 --- a/services/workers/orchestrator/src/worker/adapters/sqs_publisher.py +++ b/services/workers/orchestrator/src/worker/adapters/sqs_publisher.py @@ -97,7 +97,7 @@ async def publish(self, queue_name: str, payload: dict) -> None: self._queue_url_cache[queue_name] = resp["QueueUrl"] except Exception: logger.exception(f"Failed to get queue url for {queue_name}") - return + raise queue_url = self._queue_url_cache[queue_name] @@ -108,3 +108,4 @@ async def publish(self, queue_name: str, payload: dict) -> None: await sqs.send_message(QueueUrl=queue_url, MessageBody=json.dumps(payload)) except Exception: logger.exception(f"Failed to send single message to {queue_name}") + raise diff --git a/services/workers/orchestrator/src/worker/core/security.py b/services/workers/orchestrator/src/worker/core/security.py index 5dd9f687..90969ed2 100644 --- a/services/workers/orchestrator/src/worker/core/security.py +++ b/services/workers/orchestrator/src/worker/core/security.py @@ -1,5 +1,8 @@ import ipaddress import logging +import socket +from contextlib import contextmanager +from contextvars import ContextVar from urllib.parse import urlparse logger = logging.getLogger(__name__) @@ -50,3 +53,54 @@ def validate_target_url(url: str) -> bool: except Exception as e: logger.error(f"SSRF validation error: {e}") return False + + +_override_dns = ContextVar("override_dns", default=None) +_orig_getaddrinfo = socket.getaddrinfo + + +def _patched_getaddrinfo(host, port, family=0, type=0, proto=0, flags=0): + override = _override_dns.get() + if override and host == override[0]: + return _orig_getaddrinfo(override[1], port, family, type, proto, flags) + return _orig_getaddrinfo(host, port, family, type, proto, flags) + + +socket.getaddrinfo = _patched_getaddrinfo + + +def get_safe_ip(hostname: str) -> str | None: + try: + addr_info = socket.getaddrinfo(hostname, None) + except socket.gaierror: + return None + for addr in addr_info: + ip_str = addr[4][0] + ip = ipaddress.ip_address(ip_str) + if ip.is_private or ip.is_loopback or ip.is_link_local or ip.is_reserved or ip.is_multicast: + return None + return ip_str + return None + + +@contextmanager +def ssrf_safe_context(url: str): + """ + Context manager that pins the validated IP address for the given URL's hostname + to prevent DNS rebinding SSRF attacks. + """ + parsed = urlparse(url) + if parsed.scheme not in ("http", "https") or not parsed.hostname: + raise ValueError("Invalid URL scheme or hostname for SSRF validation") + + safe_ip = get_safe_ip(parsed.hostname) + if not safe_ip: + raise ValueError( + f"SSRF validation failed: unsafe or unresolvable hostname {parsed.hostname}" + ) + + token = _override_dns.set((parsed.hostname, safe_ip)) + try: + yield + finally: + _override_dns.reset(token) diff --git a/services/workers/orchestrator/src/worker/core/tenant_resolver.py b/services/workers/orchestrator/src/worker/core/tenant_resolver.py index 6d4be95a..fd9ad5aa 100644 --- a/services/workers/orchestrator/src/worker/core/tenant_resolver.py +++ b/services/workers/orchestrator/src/worker/core/tenant_resolver.py @@ -8,15 +8,29 @@ 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): + def __init__(self, db_router: DatabaseRouter, ttl_secs: int = 300, max_entries: int = 1000): self.db_router = db_router self._cache: dict[int, tuple[str, str, float]] = {} self._ttl = ttl_secs + self._max_entries = max_entries + + def _sweep(self, now: float) -> None: + expired = [k for k, v in self._cache.items() if v[2] <= now] + for k in expired: + del self._cache[k] + + if len(self._cache) > self._max_entries: + sorted_entries = sorted(self._cache.items(), key=lambda x: x[1][2]) + to_evict = len(self._cache) - self._max_entries + for k, _ in sorted_entries[:to_evict]: + del self._cache[k] async def resolve(self, tenant_id: int) -> tuple[str, str]: import time now = time.time() + self._sweep(now) + if tenant_id in self._cache: shard_name, shard_dsn, expiry = self._cache[tenant_id] if now < expiry: @@ -32,6 +46,7 @@ async def resolve(self, tenant_id: int) -> tuple[str, str]: 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) + self._sweep(now) return str(shard_obj.name), str(shard_obj.dsn) finally: await global_gen.aclose() diff --git a/services/workers/orchestrator/src/worker/data/handlers.py b/services/workers/orchestrator/src/worker/data/handlers.py index da829b36..cbd9a186 100644 --- a/services/workers/orchestrator/src/worker/data/handlers.py +++ b/services/workers/orchestrator/src/worker/data/handlers.py @@ -18,7 +18,7 @@ ) from pipeline.core.transformation import InboundTransformService, OutboundTransformService from worker.adapters.vault import WorkerVaultAdapter -from worker.core.security import validate_target_url +from worker.core.security import ssrf_safe_context from worker.core.tenant_resolver import TenantResolver logger = logging.getLogger(__name__) @@ -142,10 +142,10 @@ async def process_delivery( settings=get_settings(), storage=storage_adapter, ) - http_adapter = HttpxDeliveryAdapter(validator=validate_target_url) + http_adapter = HttpxDeliveryAdapter(validator=ssrf_safe_context) sftp_adapter = ParamikoSftpDeliveryAdapter() vault_adapter = WorkerVaultAdapter() - as2_adapter = HttpxAS2DeliveryAdapter() + as2_adapter = HttpxAS2DeliveryAdapter(validator=ssrf_safe_context) if idempotency_key: import uuid @@ -165,13 +165,25 @@ async def process_delivery( await session.commit() return - # Mark as processed in same transaction - session.add(ProcessedEvent(idempotency_key=key_uuid)) + # Check Outbox status + stmt = select(DataPlaneOutbox).where(DataPlaneOutbox.idempotency_key == key_uuid) + outbox_record = (await session.execute(stmt)).scalar_one_or_none() + if outbox_record: + if outbox_record.status == "DELIVERING": + logger.warning( + f"Delivery {key_uuid} is in DELIVERING state (crash/timeout). Proceeding with retry downstream..." + ) + elif outbox_record.status == "PROCESSED": + await session.commit() + return + + # Persist "DELIVERING" state await session.execute( update(DataPlaneOutbox) .where(DataPlaneOutbox.idempotency_key == key_uuid) - .values(status="PROCESSED") + .values(status="DELIVERING") ) + await session.commit() # Instantiate Domain Service strategies = { @@ -184,11 +196,32 @@ async def process_delivery( strategies=strategies, ) - # Execute pure domain logic - await service.deliver(trace_id) + try: + # Execute pure domain logic + await service.deliver( + trace_id, idempotency_key=str(key_uuid) if idempotency_key else None + ) + + if idempotency_key: + session.add(ProcessedEvent(idempotency_key=key_uuid)) + await session.execute( + update(DataPlaneOutbox) + .where(DataPlaneOutbox.idempotency_key == key_uuid) + .values(status="PROCESSED") + ) + # Commit transaction + await session.commit() + except Exception: + if idempotency_key: + await session.rollback() + await session.execute( + update(DataPlaneOutbox) + .where(DataPlaneOutbox.idempotency_key == key_uuid) + .values(status="FAILED") + ) + await session.commit() + raise - # Commit transaction - await session.commit() except Exception: await session.rollback() raise diff --git a/services/workers/orchestrator/src/worker/data/main.py b/services/workers/orchestrator/src/worker/data/main.py index d25fab5d..16d75f2a 100644 --- a/services/workers/orchestrator/src/worker/data/main.py +++ b/services/workers/orchestrator/src/worker/data/main.py @@ -121,10 +121,26 @@ async def delivery_processor(body: dict[str, Any]) -> None: ) ) - # Run the scheduler loop in the background - await scheduler_service.start(poll_interval_seconds=10.0) - - await asyncio.gather(transform_task, deliver_task, scheduled_jobs_task) + try: + # Run the scheduler loop in the background + await scheduler_service.start(poll_interval_seconds=10.0) + + await asyncio.gather(transform_task, deliver_task, scheduled_jobs_task) + finally: + logger.info("Shutting down data worker tasks gracefully...") + transform_task.cancel() + deliver_task.cancel() + scheduled_jobs_task.cancel() + + if hasattr(scheduler_service, "stop"): + await scheduler_service.stop() + + import contextlib + + with contextlib.suppress(asyncio.CancelledError): + await asyncio.gather( + transform_task, deliver_task, scheduled_jobs_task, return_exceptions=True + ) if __name__ == "__main__": diff --git a/services/workers/orchestrator/src/worker/data/scheduled_jobs_handler.py b/services/workers/orchestrator/src/worker/data/scheduled_jobs_handler.py index 3936e56d..4083cb1c 100644 --- a/services/workers/orchestrator/src/worker/data/scheduled_jobs_handler.py +++ b/services/workers/orchestrator/src/worker/data/scheduled_jobs_handler.py @@ -30,7 +30,7 @@ async def process_scheduled_job(message: dict[str, Any], **kwargs: Any) -> None: handler = registry.get(job_name) if not handler: logger.error(f"Unknown scheduled job name: {job_name}") - return + raise ValueError(f"Unknown scheduled job name: {job_name}") # Reconstruct a dummy Job object just enough for the handler to execute it job = Job(id=uuid.UUID(job_id), name=job_name, payload=job_payload) diff --git a/services/workers/orchestrator/src/worker/jobs/data_retention.py b/services/workers/orchestrator/src/worker/jobs/data_retention.py index d9b5f50f..165c4929 100644 --- a/services/workers/orchestrator/src/worker/jobs/data_retention.py +++ b/services/workers/orchestrator/src/worker/jobs/data_retention.py @@ -39,7 +39,7 @@ async def _bounded_cleanup(shard_name: str, shard_dsn: str) -> tuple[int, int]: return await self._cleanup_shard(shard_name, shard_dsn) except Exception as e: logger.error(f"[DataRetentionCleanup] Failed cleaning shard {shard_name}: {e}") - return 0, 0 + raise results = await asyncio.gather( *[_bounded_cleanup(shard.name, shard.dsn) for shard in shards] diff --git a/services/workers/orchestrator/src/worker/main.py b/services/workers/orchestrator/src/worker/main.py index c04608df..71f407a7 100644 --- a/services/workers/orchestrator/src/worker/main.py +++ b/services/workers/orchestrator/src/worker/main.py @@ -14,7 +14,16 @@ async def main() -> None: data_task = asyncio.create_task(data_main()) provision_task = asyncio.create_task(provision_main()) - await asyncio.gather(data_task, provision_task) + try: + await asyncio.gather(data_task, provision_task) + finally: + logger.info("Shutting down top-level worker tasks gracefully...") + data_task.cancel() + provision_task.cancel() + import contextlib + + with contextlib.suppress(asyncio.CancelledError): + await asyncio.gather(data_task, provision_task, return_exceptions=True) if __name__ == "__main__": diff --git a/services/workers/orchestrator/tests/test_data_main.py b/services/workers/orchestrator/tests/test_data_main.py index 33fc84d9..acb2465e 100644 --- a/services/workers/orchestrator/tests/test_data_main.py +++ b/services/workers/orchestrator/tests/test_data_main.py @@ -74,9 +74,7 @@ async def test_tenant_resolver_integration(router: DatabaseRouter): tenant_name = f"Test Tenant_{suffix}" # 1. Insert a shard and a tenant into the real DB - shard = DatabaseShard( - name=shard_name, dsn="postgresql+asyncpg://test:pass@localhost:5433/test_shard" - ) + shard = DatabaseShard(name=shard_name, dsn=SHARD_1_URL) global_session.add(shard) await global_session.commit() @@ -96,7 +94,7 @@ async def test_tenant_resolver_integration(router: DatabaseRouter): # Resolving once should hit DB resolved_shard_name, shard_dsn = await resolver.resolve(tenant_id) assert resolved_shard_name == shard_name - assert shard_dsn == "postgresql+asyncpg://test:pass@localhost:5433/test_shard" + assert shard_dsn == SHARD_1_URL # Resolving again should hit cache resolved_shard_name2, shard_dsn2 = await resolver.resolve(tenant_id) @@ -203,3 +201,27 @@ async def __aexit__(self, exc_type, exc_val, exc_tb): # Ensure all 3 messages were deleted (1 success, 2 poison) assert mock_sqs.delete_message.call_count == 3 + + +@pytest.mark.asyncio +async def test_main_execution_loop(): + from worker.data.main import main + + with ( + patch("worker.data.main.asyncio.create_task") as mock_create_task, + patch("worker.data.main.asyncio.gather", new_callable=AsyncMock) as mock_gather, + patch("worker.data.main.DatabaseRouter"), + patch("worker.data.main.TenantResolver"), + patch("worker.data.main.SqsPublisherAdapter"), + patch("worker.data.main.poll_sqs_queue", new_callable=AsyncMock), + patch("worker.data.main.SchedulerWorkerService") as mock_scheduler_cls, + ): + mock_scheduler = MagicMock() + mock_scheduler.start = AsyncMock() + mock_scheduler.stop = AsyncMock() + mock_scheduler_cls.return_value = mock_scheduler + + await main() + + assert mock_create_task.call_count >= 3 + mock_gather.assert_awaited() diff --git a/services/workers/orchestrator/tests/test_data_retention.py b/services/workers/orchestrator/tests/test_data_retention.py new file mode 100644 index 00000000..b5e09856 --- /dev/null +++ b/services/workers/orchestrator/tests/test_data_retention.py @@ -0,0 +1,91 @@ +import uuid +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from scheduler.domain.models import Job +from worker.jobs.data_retention import DataRetentionCleanupJobHandler + + +@pytest.mark.asyncio +async def test_data_retention_execute(): + db_router = MagicMock() + mock_global_session = MagicMock() + mock_global_session.execute = AsyncMock() + + mock_shard = MagicMock() + mock_shard.name = "test_shard" + mock_shard.dsn = "test_dsn" + + mock_result = MagicMock() + mock_result.scalars.return_value.all.return_value = [mock_shard] + mock_global_session.execute.return_value = mock_result + + async def get_global_session(): + yield mock_global_session + + db_router.get_global_session = get_global_session + + handler = DataRetentionCleanupJobHandler(db_router) + handler._cleanup_shard = AsyncMock(return_value=(15, 0)) + + job = Job(id=uuid.uuid4(), name="data_retention_cleanup", payload={}, interval_seconds=120) + + next_run = await handler.execute(job) + + assert handler._cleanup_shard.await_count == 1 + handler._cleanup_shard.assert_awaited_with("test_shard", "test_dsn") + + assert next_run is None + + +@pytest.mark.asyncio +async def test_data_retention_execute_exception_propagates(): + db_router = MagicMock() + mock_global_session = MagicMock() + mock_global_session.execute = AsyncMock() + + mock_shard = MagicMock() + mock_shard.name = "test_shard" + mock_shard.dsn = "test_dsn" + + mock_result = MagicMock() + mock_result.scalars.return_value.all.return_value = [mock_shard] + mock_global_session.execute.return_value = mock_result + + async def get_global_session(): + yield mock_global_session + + db_router.get_global_session = get_global_session + + handler = DataRetentionCleanupJobHandler(db_router) + handler._cleanup_shard = AsyncMock(side_effect=Exception("DB Down")) + + job = Job(id=uuid.uuid4(), name="data_retention_cleanup", payload={}, interval_seconds=120) + + with pytest.raises(Exception, match="DB Down"): + await handler.execute(job) + + +@pytest.mark.asyncio +@patch("worker.jobs.data_retention.AsyncSession") +async def test_data_retention_cleanup_shard(mock_async_session): + db_router = MagicMock() + db_router.get_engine = AsyncMock(return_value=MagicMock()) + + mock_session = MagicMock() + mock_session.execute = AsyncMock() + mock_result = MagicMock() + mock_result.rowcount = 42 + mock_session.execute.return_value = mock_result + + mock_session_ctx = MagicMock() + mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session) + mock_session_ctx.__aexit__ = AsyncMock() + + mock_async_session.return_value = mock_session_ctx + + handler = DataRetentionCleanupJobHandler(db_router) + + processed = await handler._cleanup_shard("test_shard", "test_dsn") + assert processed == (42, 42) + assert mock_session.execute.await_count == 2 diff --git a/services/workers/orchestrator/tests/test_handlers.py b/services/workers/orchestrator/tests/test_handlers.py new file mode 100644 index 00000000..8dc250c6 --- /dev/null +++ b/services/workers/orchestrator/tests/test_handlers.py @@ -0,0 +1,195 @@ +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from worker.data.handlers import process_pipeline_event + + +@pytest.mark.asyncio +@patch("worker.data.handlers.SqlAlchemyRepositoryAdapter") +@patch("worker.data.handlers.BotsTransformerAdapter") +@patch("worker.data.handlers.S3StorageAdapter") +async def test_process_pipeline_event_idempotency_duplicate(mock_s3, mock_transformer, mock_repo): + resolver = AsyncMock() + resolver.resolve.return_value = ("shard_1", "fake_dsn") + db_router = MagicMock() + + mock_session = AsyncMock() + + async def fake_get_tenant_session(*args, **kwargs): + yield mock_session + + db_router.get_tenant_session.return_value = fake_get_tenant_session() + + # We need to mock session.execute returning an existing duplicate + mock_existing = MagicMock() + mock_existing.scalar_one_or_none.return_value = MagicMock() # simulating existing record + mock_session.execute.return_value = mock_existing + + await process_pipeline_event( + trace_id="trace-123", + event_type="INBOUND", + payload={"direction": "INBOUND"}, + tenant_id=1, + resolver=resolver, + db_router=db_router, + s3_bucket="test", + aws_endpoint=None, + idempotency_key="00000000-0000-0000-0000-000000000000", + ) + # Verify it skips and returns early + mock_session.commit.assert_called_once() + mock_repo.assert_called_once() + + +@pytest.mark.asyncio +@patch("worker.data.handlers.SqlAlchemyRepositoryAdapter") +@patch("worker.data.handlers.S3StorageAdapter") +@patch("pipeline.core.saga.TraceLifecycleService") +async def test_process_pipeline_event_transform_completed(mock_trace, mock_s3, mock_repo): + resolver = AsyncMock() + resolver.resolve.return_value = ("shard_1", "fake_dsn") + db_router = MagicMock() + + mock_session = AsyncMock() + + async def fake_get_tenant_session(*args, **kwargs): + yield mock_session + + db_router.get_tenant_session.return_value = fake_get_tenant_session() + + # We mock it so it does not find a duplicate + mock_existing = MagicMock() + mock_existing.scalar_one_or_none.return_value = None + mock_session.execute.return_value = mock_existing + + mock_saga = AsyncMock() + mock_trace.return_value = mock_saga + + await process_pipeline_event( + trace_id="trace-123", + event_type="TRANSFORM_COMPLETED", + payload={"direction": "INBOUND"}, + tenant_id=1, + resolver=resolver, + db_router=db_router, + s3_bucket="test", + aws_endpoint=None, + idempotency_key="00000000-0000-0000-0000-000000000000", + ) + mock_saga.handle_transform_completed.assert_called_once_with({"direction": "INBOUND"}) + + +@pytest.mark.asyncio +@patch("worker.data.handlers.SqlAlchemyRepositoryAdapter") +@patch("worker.data.handlers.S3StorageAdapter") +@patch("pipeline.core.saga.TraceLifecycleService") +async def test_process_pipeline_event_delivery_completed(mock_trace, mock_s3, mock_repo): + resolver = AsyncMock() + resolver.resolve.return_value = ("shard_1", "fake_dsn") + db_router = MagicMock() + + mock_session = AsyncMock() + + async def fake_get_tenant_session(*args, **kwargs): + yield mock_session + + db_router.get_tenant_session.return_value = fake_get_tenant_session() + + # We mock it so it does not find a duplicate + mock_existing = MagicMock() + mock_existing.scalar_one_or_none.return_value = None + mock_session.execute.return_value = mock_existing + + mock_saga = AsyncMock() + mock_trace.return_value = mock_saga + + await process_pipeline_event( + trace_id="trace-123", + event_type="DELIVERY_COMPLETED", + payload={"direction": "INBOUND"}, + tenant_id=1, + resolver=resolver, + db_router=db_router, + s3_bucket="test", + aws_endpoint=None, + idempotency_key="00000000-0000-0000-0000-000000000000", + ) + mock_saga.handle_delivery_completed.assert_called_once_with({"direction": "INBOUND"}) + + +@pytest.mark.asyncio +@patch("worker.data.handlers.SqlAlchemyRepositoryAdapter") +@patch("worker.data.handlers.S3StorageAdapter") +@patch("worker.data.handlers.InboundTransformService") +async def test_process_pipeline_event_inbound(mock_inbound, mock_s3, mock_repo): + resolver = AsyncMock() + resolver.resolve.return_value = ("shard_1", "fake_dsn") + db_router = MagicMock() + + mock_session = AsyncMock() + + async def fake_get_tenant_session(*args, **kwargs): + yield mock_session + + db_router.get_tenant_session.return_value = fake_get_tenant_session() + + # We mock it so it does not find a duplicate + mock_existing = MagicMock() + mock_existing.scalar_one_or_none.return_value = None + mock_session.execute.return_value = mock_existing + + mock_service = AsyncMock() + mock_inbound.return_value = mock_service + + await process_pipeline_event( + trace_id="trace-123", + event_type="DOCUMENT_RECEIVED", + payload={"direction": "INBOUND"}, + tenant_id=1, + resolver=resolver, + db_router=db_router, + s3_bucket="test", + aws_endpoint=None, + idempotency_key="00000000-0000-0000-0000-000000000000", + ) + mock_inbound.assert_called_once() + mock_service.transform.assert_called_once_with("trace-123") + + +@pytest.mark.asyncio +@patch("worker.data.handlers.SqlAlchemyRepositoryAdapter") +@patch("worker.data.handlers.S3StorageAdapter") +@patch("worker.data.handlers.OutboundTransformService") +async def test_process_pipeline_event_outbound(mock_outbound, mock_s3, mock_repo): + resolver = AsyncMock() + resolver.resolve.return_value = ("shard_1", "fake_dsn") + db_router = MagicMock() + + mock_session = AsyncMock() + + async def fake_get_tenant_session(*args, **kwargs): + yield mock_session + + db_router.get_tenant_session.return_value = fake_get_tenant_session() + + # We mock it so it does not find a duplicate + mock_existing = MagicMock() + mock_existing.scalar_one_or_none.return_value = None + mock_session.execute.return_value = mock_existing + + mock_service = AsyncMock() + mock_outbound.return_value = mock_service + + await process_pipeline_event( + trace_id="trace-123", + event_type="OUTBOUND_REQUEST", + payload={"direction": "OUTBOUND"}, + tenant_id=1, + resolver=resolver, + db_router=db_router, + s3_bucket="test", + aws_endpoint=None, + idempotency_key="00000000-0000-0000-0000-000000000000", + ) + mock_outbound.assert_called_once() + mock_service.transform.assert_called_once_with("trace-123") diff --git a/services/workers/orchestrator/tests/test_outbox_sweeper.py b/services/workers/orchestrator/tests/test_outbox_sweeper.py new file mode 100644 index 00000000..489d2557 --- /dev/null +++ b/services/workers/orchestrator/tests/test_outbox_sweeper.py @@ -0,0 +1,141 @@ +import datetime +import uuid +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from scheduler.domain.models import Job +from worker.jobs.outbox_sweeper import DataPlaneOutboxSweeperJobHandler + + +@pytest.mark.asyncio +async def test_outbox_sweeper_execute(): + db_router = MagicMock() + message_publisher = MagicMock() + message_publisher.connect.return_value.__aenter__ = AsyncMock() + message_publisher.connect.return_value.__aexit__ = AsyncMock() + + mock_global_session = MagicMock() + mock_global_session.execute = AsyncMock() + + mock_shard = MagicMock() + mock_shard.name = "test_shard" + mock_shard.dsn = "test_dsn" + + mock_result = MagicMock() + mock_result.scalars.return_value.all.return_value = [mock_shard] + mock_global_session.execute.return_value = mock_result + + async def get_global_session(): + yield mock_global_session + + db_router.get_global_session = get_global_session + + handler = DataPlaneOutboxSweeperJobHandler(db_router, message_publisher) + handler._sweep_shard = AsyncMock(return_value=5) + + job = Job(id=uuid.uuid4(), name="outbox_sweeper", payload={}, interval_seconds=120) + + next_run = await handler.execute(job) + + assert handler._sweep_shard.await_count == 1 + handler._sweep_shard.assert_awaited_with("test_shard", "test_dsn") + + assert next_run is not None + assert isinstance(next_run, datetime.datetime) + + +@pytest.mark.asyncio +async def test_outbox_sweeper_execute_exception_caught(): + db_router = MagicMock() + message_publisher = MagicMock() + message_publisher.connect.return_value.__aenter__ = AsyncMock() + message_publisher.connect.return_value.__aexit__ = AsyncMock() + + mock_global_session = MagicMock() + mock_global_session.execute = AsyncMock() + + mock_shard = MagicMock() + mock_shard.name = "test_shard" + mock_shard.dsn = "test_dsn" + + mock_result = MagicMock() + mock_result.scalars.return_value.all.return_value = [mock_shard] + mock_global_session.execute.return_value = mock_result + + async def get_global_session(): + yield mock_global_session + + db_router.get_global_session = get_global_session + + handler = DataPlaneOutboxSweeperJobHandler(db_router, message_publisher) + handler._sweep_shard = AsyncMock(side_effect=Exception("Database down")) + + job = Job(id=uuid.uuid4(), name="outbox_sweeper", payload={}, interval_seconds=120) + + next_run = await handler.execute(job) + + assert handler._sweep_shard.await_count == 1 + assert next_run is not None + + +@pytest.mark.asyncio +@patch("worker.jobs.outbox_sweeper.AsyncSession") +async def test_outbox_sweep_shard_no_events(mock_async_session): + db_router = MagicMock() + db_router.get_engine = AsyncMock(return_value=MagicMock()) + message_publisher = MagicMock() + + mock_session = MagicMock() + mock_session.execute = AsyncMock() + mock_result = MagicMock() + mock_result.scalars.return_value.all.return_value = [] + mock_session.execute.return_value = mock_result + + mock_session_ctx = MagicMock() + mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session) + mock_session_ctx.__aexit__ = AsyncMock() + + mock_async_session.return_value = mock_session_ctx + + handler = DataPlaneOutboxSweeperJobHandler(db_router, message_publisher) + + processed = await handler._sweep_shard("test_shard", "test_dsn") + assert processed == 0 + + +@pytest.mark.asyncio +@patch("worker.jobs.outbox_sweeper.AsyncSession") +async def test_outbox_sweep_shard_with_events(mock_async_session): + db_router = MagicMock() + db_router.get_engine = AsyncMock(return_value=MagicMock()) + message_publisher = MagicMock() + message_publisher.publish_batch = AsyncMock(return_value=["1", "2"]) + + mock_session = MagicMock() + mock_session.execute = AsyncMock() + + mock_event1 = MagicMock() + mock_event1.id = "1" + mock_event1.event_type = "TRANSFORM_EVENT" + mock_event1.payload = {"foo": "bar"} + + mock_event2 = MagicMock() + mock_event2.id = "2" + mock_event2.event_type = "DELIVER_EVENT" + mock_event2.payload = {"baz": "qux"} + + mock_result = MagicMock() + mock_result.scalars.return_value.all.return_value = [mock_event1, mock_event2] + mock_session.execute.return_value = mock_result + + mock_session_ctx = MagicMock() + mock_session_ctx.__aenter__ = AsyncMock(return_value=mock_session) + mock_session_ctx.__aexit__ = AsyncMock() + + mock_async_session.return_value = mock_session_ctx + + handler = DataPlaneOutboxSweeperJobHandler(db_router, message_publisher) + + processed = await handler._sweep_shard("test_shard", "test_dsn") + assert processed == 2 + message_publisher.publish_batch.assert_awaited() diff --git a/services/workers/orchestrator/tests/test_scheduled_jobs_handler.py b/services/workers/orchestrator/tests/test_scheduled_jobs_handler.py new file mode 100644 index 00000000..891214df --- /dev/null +++ b/services/workers/orchestrator/tests/test_scheduled_jobs_handler.py @@ -0,0 +1,51 @@ +import uuid +from unittest.mock import AsyncMock, MagicMock + +import pytest +from worker.data.scheduled_jobs_handler import process_scheduled_job + + +@pytest.mark.asyncio +async def test_process_scheduled_job_missing_id(): + # Should log warning and return + await process_scheduled_job({"job_name": "test"}) + + +@pytest.mark.asyncio +async def test_process_scheduled_job_missing_name(): + # Should log warning and return + await process_scheduled_job({"job_id": str(uuid.uuid4())}) + + +@pytest.mark.asyncio +async def test_process_scheduled_job_missing_registry(): + # Should log error and return + await process_scheduled_job({"job_id": str(uuid.uuid4()), "job_name": "test_job"}) + + +@pytest.mark.asyncio +async def test_process_scheduled_job_unknown_job(): + registry = MagicMock() + registry.get.return_value = None + with pytest.raises(ValueError, match="Unknown scheduled job name: unknown_job"): + await process_scheduled_job( + {"job_id": str(uuid.uuid4()), "job_name": "unknown_job"}, registry=registry + ) + + +@pytest.mark.asyncio +async def test_process_scheduled_job_success(): + registry = MagicMock() + mock_handler = AsyncMock() + registry.get.return_value = mock_handler + + job_id = str(uuid.uuid4()) + await process_scheduled_job( + {"job_id": job_id, "job_name": "known_job", "payload": {"foo": "bar"}}, registry=registry + ) + + mock_handler.execute.assert_awaited_once() + job = mock_handler.execute.call_args[0][0] + assert str(job.id) == job_id + assert job.name == "known_job" + assert job.payload == {"foo": "bar"} diff --git a/services/workers/orchestrator/tests/test_security.py b/services/workers/orchestrator/tests/test_security.py new file mode 100644 index 00000000..62aa3820 --- /dev/null +++ b/services/workers/orchestrator/tests/test_security.py @@ -0,0 +1,58 @@ +from unittest.mock import patch + +import pytest +from worker.core.security import get_safe_ip, ssrf_safe_context, validate_target_url + + +def test_validate_target_url_invalid_scheme(): + assert not validate_target_url("ftp://example.com") + assert not validate_target_url("file:///etc/passwd") + + +def test_validate_target_url_no_hostname(): + assert not validate_target_url("http://") + + +@patch("socket.getaddrinfo") +def test_validate_target_url_valid(mock_getaddrinfo): + mock_getaddrinfo.return_value = [(2, 1, 6, "", ("93.184.216.34", 80))] + assert validate_target_url("http://example.com") + + +@patch("socket.getaddrinfo") +def test_validate_target_url_private_ip(mock_getaddrinfo): + mock_getaddrinfo.return_value = [(2, 1, 6, "", ("192.168.1.1", 80))] + assert not validate_target_url("http://internal.com") + + +@patch("socket.getaddrinfo") +def test_validate_target_url_loopback_ip(mock_getaddrinfo): + mock_getaddrinfo.return_value = [(2, 1, 6, "", ("127.0.0.1", 80))] + assert not validate_target_url("http://localhost") + + +@patch("socket.getaddrinfo") +def test_get_safe_ip(mock_getaddrinfo): + mock_getaddrinfo.return_value = [(2, 1, 6, "", ("93.184.216.34", 80))] + assert get_safe_ip("example.com") == "93.184.216.34" + + +@patch("socket.getaddrinfo") +def test_get_safe_ip_private(mock_getaddrinfo): + mock_getaddrinfo.return_value = [(2, 1, 6, "", ("192.168.1.1", 80))] + assert get_safe_ip("example.com") is None + + +@patch("socket.getaddrinfo") +def test_ssrf_safe_context_valid(mock_getaddrinfo): + mock_getaddrinfo.return_value = [(2, 1, 6, "", ("93.184.216.34", 80))] + with ssrf_safe_context("http://example.com"): + import socket + + res = socket.getaddrinfo("example.com", 80) + assert res == [(2, 1, 6, "", ("93.184.216.34", 80))] + + +def test_ssrf_safe_context_invalid_url(): + with pytest.raises(ValueError), ssrf_safe_context("ftp://example.com"): + pass From c6c6cd1b3cc118206fb05a2fdf0da3d99569d345 Mon Sep 17 00:00:00 2001 From: Pramod Date: Fri, 17 Jul 2026 22:29:33 +0530 Subject: [PATCH 5/5] light weight scheduler and outbox sweeper code review --- .../components/CreatePartnerModal.tsx | 14 ++- libs/as2_core/src/as2_core/builder.py | 6 +- .../src/pipeline/core/as2_orchestrator.py | 2 + .../src/pipeline/core/delivery/as2.py | 17 ++-- .../src/pipeline/core/delivery/router.py | 13 ++- libs/pipeline/tests/fakes.py | 17 ++++ libs/pipeline/tests/test_delivery_service.py | 1 + .../tests/test_delivery_service_as2.py | 16 +++- .../api/src/api/routers/platform/scheduler.py | 14 +++ .../trading_partners/platform/as2_partners.py | 5 +- .../api/tests/test_api_receiver_service.py | 22 ++++- .../src/worker/core/tenant_resolver.py | 2 +- .../orchestrator/tests/test_security.py | 3 +- .../tests/test_tenant_resolver.py | 86 +++++++++++++++++++ 14 files changed, 197 insertions(+), 21 deletions(-) create mode 100644 services/workers/orchestrator/tests/test_tenant_resolver.py diff --git a/frontend/web/src/features/partners/components/CreatePartnerModal.tsx b/frontend/web/src/features/partners/components/CreatePartnerModal.tsx index 7f05216d..543456cf 100644 --- a/frontend/web/src/features/partners/components/CreatePartnerModal.tsx +++ b/frontend/web/src/features/partners/components/CreatePartnerModal.tsx @@ -1,4 +1,4 @@ -import { useState } from 'react'; +import { useState, useRef, useEffect } from 'react'; import { Input } from '@/components/ui/input'; import { Label } from '@/components/ui/label'; import { FormModal } from '@/components/ui/form-modal'; @@ -20,6 +20,12 @@ export function CreatePartnerModal({ existingAs2Ids = [] }: { existingAs2Ids?: s const [as2Id, setAs2Id] = useState(''); const [url, setUrl] = useState(''); + const isOpenRef = useRef(isOpen); + useEffect(() => { isOpenRef.current = isOpen; }, [isOpen]); + + const isLocalRef = useRef(isLocal); + useEffect(() => { isLocalRef.current = isLocal; }, [isLocal]); + const isDuplicate = existingAs2Ids.includes(as2Id); const { data: platformSettings } = usePlatformSettings(); @@ -272,6 +278,12 @@ export function CreatePartnerModal({ existingAs2Ids = [] }: { existingAs2Ids?: s } generateCert.mutate(as2Id, { onSuccess: (res) => { + if (!isOpenRef.current || !isLocalRef.current) { + // The modal was closed or switched to remote while the mutation was inflight. + // Cleanup the newly created orphaned secret immediately. + deleteCertSecret.mutate(res.private_key_vault_ref); + return; + } setCertPem(res.public_cert_pem); setPrivateKeyVaultRef(res.private_key_vault_ref); setGeneratedForAs2Id(as2Id); diff --git a/libs/as2_core/src/as2_core/builder.py b/libs/as2_core/src/as2_core/builder.py index 426f78eb..72b3f256 100644 --- a/libs/as2_core/src/as2_core/builder.py +++ b/libs/as2_core/src/as2_core/builder.py @@ -52,6 +52,7 @@ def build_outbound_message( encrypt_fn: object | None = None, mdn_url: str | None = None, mic_alg: str = "sha256", + message_id: str | None = None, ) -> OutboundAS2Message: """ Builds a fully-wrapped, optionally signed, optionally encrypted AS2 @@ -131,13 +132,12 @@ def build_outbound_message( is_encrypted = True - # ── Step 5: Build HTTP Headers ──────────────────────────────────────────── - message_id = f"<{uuid.uuid4()}@soopaedi>" + message_id_str = f"<{message_id}@soopaedi>" if message_id else f"<{uuid.uuid4()}@soopaedi>" headers: dict[str, str] = { "AS2-Version": "1.2", "AS2-From": as2_from, "AS2-To": as2_to, - "Message-ID": message_id, + "Message-ID": message_id_str, "Content-Type": current_content_type, "MIME-Version": "1.0", "Disposition-Notification-To": as2_from, diff --git a/libs/pipeline/src/pipeline/core/as2_orchestrator.py b/libs/pipeline/src/pipeline/core/as2_orchestrator.py index ea1ac423..fcec0123 100644 --- a/libs/pipeline/src/pipeline/core/as2_orchestrator.py +++ b/libs/pipeline/src/pipeline/core/as2_orchestrator.py @@ -59,6 +59,7 @@ async def build( raw_payload: bytes, local_partner: dict[str, Any] | None, remote_partner: dict[str, Any], + idempotency_key: str | None = None, ) -> OutboundAS2Message: """ Builds the fully-wrapped AS2 HTTP message for transmission. @@ -136,4 +137,5 @@ async def build( sign_fn=sign_fn, encrypt_fn=encrypt_fn, mdn_url=mdn_url, + message_id=idempotency_key, ) diff --git a/libs/pipeline/src/pipeline/core/delivery/as2.py b/libs/pipeline/src/pipeline/core/delivery/as2.py index 0c720169..80534ad3 100644 --- a/libs/pipeline/src/pipeline/core/delivery/as2.py +++ b/libs/pipeline/src/pipeline/core/delivery/as2.py @@ -57,14 +57,17 @@ async def deliver( raw_payload=raw_payload, local_partner=local_partner, remote_partner=remote_partner, + idempotency_key=idempotency_key, ) - except Exception: + except Exception as e: await self.repository.update_edi_message_status(trace_id, MessageStatus.FAILED) await self._emit_delivery_completed(trace_id, edi_msg.direction, MessageStatus.FAILED) logger.exception( f"AS2 Delivery Adapter is misconfigured or failed to build for trace_id={trace_id}" ) - return + raise RuntimeError( + f"AS2 Delivery Adapter failed to build for trace_id={trace_id}" + ) from e try: status_code, response_headers, response_body = await self.as2_delivery.deliver( @@ -72,16 +75,18 @@ async def deliver( body=as2_msg.body, headers=as2_msg.headers, ) - except RuntimeError: + except RuntimeError as e: await self.repository.update_edi_message_status(trace_id, MessageStatus.FAILED) await self._emit_delivery_completed(trace_id, edi_msg.direction, MessageStatus.FAILED) logger.exception(f"AS2 Delivery Adapter is misconfigured for trace_id={trace_id}") - return - except Exception: + raise RuntimeError( + f"AS2 Delivery Adapter is misconfigured for trace_id={trace_id}" + ) from e + except Exception as e: await self.repository.update_edi_message_status(trace_id, MessageStatus.FAILED) await self._emit_delivery_completed(trace_id, edi_msg.direction, MessageStatus.FAILED) logger.exception(f"AS2 HTTP transmission failed for trace_id={trace_id}") - return + raise RuntimeError(f"AS2 HTTP transmission failed for trace_id={trace_id}") from e if 200 <= status_code < 300: from as2_core import parse_mdn diff --git a/libs/pipeline/src/pipeline/core/delivery/router.py b/libs/pipeline/src/pipeline/core/delivery/router.py index 93c1037f..c8cfb409 100644 --- a/libs/pipeline/src/pipeline/core/delivery/router.py +++ b/libs/pipeline/src/pipeline/core/delivery/router.py @@ -32,7 +32,12 @@ async def deliver(self, trace_id: str, idempotency_key: str | None = None) -> No direction = edi_msg.direction - if direction == "OUTBOUND" and edi_msg.trading_partner_id: + if direction == "OUTBOUND": + if not edi_msg.trading_partner_id: + raise ValueError( + f"EDI Message {trace_id} is missing trading_partner_id for OUTBOUND routing." + ) + route = await self.repository.get_outbound_route_by_trading_partner_id( trading_partner_id=edi_msg.trading_partner_id, tenant_id=edi_msg.tenant_id, @@ -61,11 +66,13 @@ async def deliver(self, trace_id: str, idempotency_key: str | None = None) -> No 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, strategy in self.strategies.items(): partner_id = route.get(route_key) if partner_id: - await strategy.deliver(trace_id, partner_id, edi_msg, idempotency_key) + try: + await strategy.deliver(trace_id, partner_id, edi_msg, idempotency_key) + except Exception as e: + raise RuntimeError(f"Delivery strategy failed for trace_id={trace_id}") from e return raise ValueError( diff --git a/libs/pipeline/tests/fakes.py b/libs/pipeline/tests/fakes.py index 59731df4..08b4349b 100644 --- a/libs/pipeline/tests/fakes.py +++ b/libs/pipeline/tests/fakes.py @@ -207,6 +207,23 @@ async def get_route( wildcard_match = next((r for r in candidates if r.get("transaction_type") == "*"), None) return wildcard_match + async def get_outbound_route_by_trading_partner_id( + self, trading_partner_id: str, tenant_id: int + ) -> dict[str, Any] | None: + candidates = [ + r + for r in self.routes + if r.get("direction") == "OUTBOUND" + and ( + r.get("sftp_partner_id") == trading_partner_id + or r.get("as2_partner_id") == trading_partner_id + or r.get("webhook_partner_id") == trading_partner_id + ) + ] + if candidates: + return candidates[0] + return None + async def get_sftp_partner(self, partner_id: str) -> dict[str, Any] | None: return self.sftp_partners.get(partner_id) diff --git a/libs/pipeline/tests/test_delivery_service.py b/libs/pipeline/tests/test_delivery_service.py index 0b373d14..4f7a7a19 100644 --- a/libs/pipeline/tests/test_delivery_service.py +++ b/libs/pipeline/tests/test_delivery_service.py @@ -114,6 +114,7 @@ async def test_delivery_service_outbound_sftp() -> None: "direction": MessageDirection.OUTBOUND, "sender_id": "SENDER1", "receiver_id": "RECV1", + "trading_partner_id": "sftp1", "transaction_type": "855", "edi_data": "FAKE*EDI*DATA~", "status": MessageStatus.PENDING_DELIVERY, diff --git a/libs/pipeline/tests/test_delivery_service_as2.py b/libs/pipeline/tests/test_delivery_service_as2.py index 0f2267dd..f54010ba 100644 --- a/libs/pipeline/tests/test_delivery_service_as2.py +++ b/libs/pipeline/tests/test_delivery_service_as2.py @@ -81,6 +81,7 @@ def _seed_as2_route( "direction": "OUTBOUND", "sender_id": "SENDER", "receiver_id": "RECEIVER", + "trading_partner_id": partner_id, "transaction_type": transaction_type, "edi_data": edi_data, "status": "PENDING_DELIVERY", @@ -156,6 +157,7 @@ async def test_deliver_as2_http_failure_sets_failed_status() -> None: "direction": "OUTBOUND", "sender_id": "S1", "receiver_id": "R1", + "trading_partner_id": "p-fail", "transaction_type": "856", "edi_data": "FAKE*EDI~", "status": "PENDING_DELIVERY", @@ -202,9 +204,12 @@ async def test_deliver_as2_null_adapter_is_caught_and_marked_failed() -> None: # ── Act / Assert ─────────────────────────────────────────────────────────── service = make_service(repo=repo, as2=NullAS2DeliveryAdapter()) - await service.deliver(trace_id) + # It should catch the RuntimeError and mark the message as FAILED, but bubble up the error + import pytest + + with pytest.raises(RuntimeError): + await service.deliver(trace_id) - # It should catch the RuntimeError and mark the message as FAILED assert len(repo.outbox) == 1 outbox_event = repo.outbox[0] assert outbox_event["event_type"] == PipelineEventType.DELIVERY_COMPLETED @@ -231,6 +236,7 @@ async def test_deliver_as2_idempotent_claim() -> None: "direction": "OUTBOUND", "sender_id": "A", "receiver_id": "B", + "trading_partner_id": "p-idem", "transaction_type": "810", "edi_data": "EDI~", "status": "PROCESSING", @@ -276,6 +282,7 @@ async def test_deliver_as2_missing_local_partner_sets_failed() -> None: "direction": "OUTBOUND", "sender_id": "X", "receiver_id": "Y", + "trading_partner_id": "p-nolocal", "transaction_type": "850", "edi_data": "EDI~", "status": "PENDING_DELIVERY", @@ -295,7 +302,10 @@ async def test_deliver_as2_missing_local_partner_sets_failed() -> None: # Do NOT seed local_as2_partners["missing-local"] # ── Act ──────────────────────────────────────────────────────────────────── - await make_service(repo=repo, as2=as2_adapter).deliver(trace_id) + import pytest + + with pytest.raises(RuntimeError): + await make_service(repo=repo, as2=as2_adapter).deliver(trace_id) # ── Assert ───────────────────────────────────────────────────────────────── assert len(repo.outbox) == 1 diff --git a/services/api/src/api/routers/platform/scheduler.py b/services/api/src/api/routers/platform/scheduler.py index 0ab70487..b0f050c6 100644 --- a/services/api/src/api/routers/platform/scheduler.py +++ b/services/api/src/api/routers/platform/scheduler.py @@ -90,6 +90,20 @@ async def create_job(request: JobCreateRequest, uow: UnitOfWork = Depends(get_uo raise HTTPException( status_code=422, detail=f"Job '{request.name}' has no configured target queue." ) + if ( + job_def.min_interval_seconds is not None + and request.interval_seconds < job_def.min_interval_seconds + ): + raise HTTPException( + status_code=422, + detail=f"Interval ({request.interval_seconds}s) is below minimum allowed ({job_def.min_interval_seconds}s).", + ) + + if job_def.max_interval_seconds and request.interval_seconds > job_def.max_interval_seconds: + raise HTTPException( + status_code=422, + detail=f"Interval ({request.interval_seconds}s) exceeds maximum allowed ({job_def.max_interval_seconds}s).", + ) async with uow: try: 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 ee9a3b7a..bcffc05c 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 @@ -99,6 +99,7 @@ async def create_platform_as2_partner( public_cert_pem = request.public_cert_pem private_key_vault_ref = request.private_key_vault_ref auto_generated = False + commit_success = False try: async with uow: @@ -140,6 +141,8 @@ async def create_platform_as2_partner( entity = await svc.create_as2_partner(tenant_id=0, cmd=cmd) await uow.commit() + commit_success = True + p = await uow.as2_partners.get_as2_partner(tenant_id=0, partner_id=entity.partner_id) if not p: raise HTTPException(status_code=500, detail="Partner creation failed") @@ -153,7 +156,7 @@ async def create_platform_as2_partner( active=p.active, ) except Exception as e: - if auto_generated and private_key_vault_ref: + if auto_generated and private_key_vault_ref and not commit_success: vault.delete_secret(private_key_vault_ref) if isinstance(e, IntegrityError): raise HTTPException( diff --git a/services/api/tests/test_api_receiver_service.py b/services/api/tests/test_api_receiver_service.py index a2c463d7..1379837f 100644 --- a/services/api/tests/test_api_receiver_service.py +++ b/services/api/tests/test_api_receiver_service.py @@ -16,7 +16,11 @@ async def test_process_api_edi_json_success(): ) assert trace_id is not None + mock_uow.transactions.create_edi_json.assert_awaited_once() + create_args, create_kwargs = mock_uow.transactions.create_edi_json.await_args + assert create_kwargs["payload"]["transaction_type"] == "850" + mock_uow.data_plane_outbox.publish_outbox_event.assert_awaited_once() args, kwargs = mock_uow.data_plane_outbox.publish_outbox_event.call_args @@ -39,6 +43,8 @@ async def test_process_api_edi_json_heading(): ) assert trace_id is not None mock_uow.transactions.create_edi_json.assert_awaited_once() + create_args, create_kwargs = mock_uow.transactions.create_edi_json.await_args + assert create_kwargs["payload"]["transaction_type"] == "850" @pytest.mark.asyncio @@ -50,16 +56,28 @@ async def test_process_api_edi_json_st_segment(): ) assert trace_id is not None mock_uow.transactions.create_edi_json.assert_awaited_once() + create_args, create_kwargs = mock_uow.transactions.create_edi_json.await_args + assert create_kwargs["payload"]["transaction_type"] == "855" @pytest.mark.asyncio async def test_process_api_edi_json_list_extraction(): mock_uow = AsyncMock() svc = ApiReceiverService(mock_uow) - # Give a list payload to hit the extraction logic for lists - payload = [{"transaction_type": "850", "foo": "bar"}, {"transaction_type": "850", "foo": "baz"}] + payload = [ + {"ST": {"ST01": "850"}, "BEG": {"BEG03": "123"}, "foo": "bar"}, + {"ST": {"ST01": "850"}, "BEG": {"BEG03": "456"}, "foo": "baz"}, + ] trace_id = await svc.process_api_edi_json( tenant_id=1, trading_partner_id="PARTNER_X", payload=payload ) assert trace_id is not None mock_uow.transactions.create_edi_json.assert_awaited_once() + create_args, create_kwargs = mock_uow.transactions.create_edi_json.await_args + + # Assert business_metadata aggregation for lists + assert create_kwargs["payload"]["business_metadata"] == { + "po_number": ["123", "456"], + "business_reference": ["123", "456"], + "_routing": {"trading_partner_id": "PARTNER_X"}, + } diff --git a/services/workers/orchestrator/src/worker/core/tenant_resolver.py b/services/workers/orchestrator/src/worker/core/tenant_resolver.py index fd9ad5aa..5418978f 100644 --- a/services/workers/orchestrator/src/worker/core/tenant_resolver.py +++ b/services/workers/orchestrator/src/worker/core/tenant_resolver.py @@ -28,7 +28,7 @@ def _sweep(self, now: float) -> None: async def resolve(self, tenant_id: int) -> tuple[str, str]: import time - now = time.time() + now = time.monotonic() self._sweep(now) if tenant_id in self._cache: diff --git a/services/workers/orchestrator/tests/test_security.py b/services/workers/orchestrator/tests/test_security.py index 62aa3820..98333e16 100644 --- a/services/workers/orchestrator/tests/test_security.py +++ b/services/workers/orchestrator/tests/test_security.py @@ -43,7 +43,7 @@ def test_get_safe_ip_private(mock_getaddrinfo): assert get_safe_ip("example.com") is None -@patch("socket.getaddrinfo") +@patch("worker.core.security._orig_getaddrinfo") def test_ssrf_safe_context_valid(mock_getaddrinfo): mock_getaddrinfo.return_value = [(2, 1, 6, "", ("93.184.216.34", 80))] with ssrf_safe_context("http://example.com"): @@ -51,6 +51,7 @@ def test_ssrf_safe_context_valid(mock_getaddrinfo): res = socket.getaddrinfo("example.com", 80) assert res == [(2, 1, 6, "", ("93.184.216.34", 80))] + mock_getaddrinfo.assert_called_with("93.184.216.34", 80, 0, 0, 0, 0) def test_ssrf_safe_context_invalid_url(): diff --git a/services/workers/orchestrator/tests/test_tenant_resolver.py b/services/workers/orchestrator/tests/test_tenant_resolver.py new file mode 100644 index 00000000..bcab7bd4 --- /dev/null +++ b/services/workers/orchestrator/tests/test_tenant_resolver.py @@ -0,0 +1,86 @@ +from unittest.mock import AsyncMock, MagicMock + +import pytest +from worker.core.tenant_resolver import TenantResolver + + +@pytest.mark.asyncio +async def test_tenant_resolver_success(): + mock_db_router = MagicMock() + mock_global_session = AsyncMock() + mock_db_router.get_global_session.return_value = mock_global_session + + class MockRow: + def __init__(self, name, dsn): + self.name = name + self.dsn = dsn + + class MockResult: + def first(self): + return (None, MockRow("shard_1", "postgresql://user:pass@host/db")) + + mock_global_session.__anext__.return_value = mock_global_session + mock_global_session.execute.return_value = MockResult() + + resolver = TenantResolver(db_router=mock_db_router, ttl_secs=300) + + # First resolve should hit DB + shard_name, shard_dsn = await resolver.resolve(tenant_id=1) + assert shard_name == "shard_1" + assert shard_dsn == "postgresql://user:pass@host/db" + mock_global_session.execute.assert_awaited_once() + + # Second resolve should hit cache + mock_global_session.execute.reset_mock() + shard_name_2, shard_dsn_2 = await resolver.resolve(tenant_id=1) + assert shard_name_2 == "shard_1" + assert shard_dsn_2 == "postgresql://user:pass@host/db" + mock_global_session.execute.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_tenant_resolver_not_found(): + mock_db_router = MagicMock() + mock_global_session = AsyncMock() + mock_db_router.get_global_session.return_value = mock_global_session + + class MockEmptyResult: + def first(self): + return None + + mock_global_session.__anext__.return_value = mock_global_session + mock_global_session.execute.return_value = MockEmptyResult() + + resolver = TenantResolver(db_router=mock_db_router, ttl_secs=300) + + with pytest.raises(ValueError, match="Tenant 999 not found in Global DB"): + await resolver.resolve(tenant_id=999) + + +@pytest.mark.asyncio +async def test_tenant_resolver_eviction(): + mock_db_router = MagicMock() + mock_global_session = AsyncMock() + mock_db_router.get_global_session.return_value = mock_global_session + + class MockRow: + def __init__(self, name, dsn): + self.name = name + self.dsn = dsn + + class MockResult: + def first(self): + return (None, MockRow("shard_1", "postgresql://user:pass@host/db")) + + mock_global_session.__anext__.return_value = mock_global_session + mock_global_session.execute.return_value = MockResult() + + # Small cache size to force eviction + resolver = TenantResolver(db_router=mock_db_router, ttl_secs=300, max_entries=2) + + await resolver.resolve(tenant_id=1) + await resolver.resolve(tenant_id=2) + # This should evict tenant 1 or 2 + await resolver.resolve(tenant_id=3) + + assert len(resolver._cache) == 2