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..0fb609be 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) @@ -25,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/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..6c3d589f 100644 --- a/frontend/web/src/features/partners/api/IPartnersRepository.ts +++ b/frontend/web/src/features/partners/api/IPartnersRepository.ts @@ -17,6 +17,7 @@ import type { */ export interface IPartnersRepository { // Platform Trading Partners + deleteCertificateSecret(vaultRef: string): Promise; getPlatformPartners(): Promise; createPlatformPartner(payload: CreatePartnerPayload): Promise; updatePlatformPartner(id: string, payload: UpdatePartnerPayload): Promise; @@ -32,6 +33,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..b0097a64 100644 --- a/frontend/web/src/features/partners/api/partnerHooks.ts +++ b/frontend/web/src/features/partners/api/partnerHooks.ts @@ -127,9 +127,8 @@ export function useUpdatePlatformPartnerMutation() { ); } -export function useDeletePlatformPartner() { +export function useDeletePlatformPartnerMutation() { const repo = useRepository(); - return useToastMutation( (id: string) => repo.deletePlatformPartner(id), 'Partner deleted.', @@ -137,6 +136,13 @@ export function useDeletePlatformPartner() { ); } +export function useDeleteCertificateSecretMutation() { + const repo = useRepository(); + return useMutation({ + mutationFn: (vaultRef: string) => repo.deleteCertificateSecret(vaultRef), + }); +} + // ───────────────────────────────────────────── // Platform Partnership Mutations // ───────────────────────────────────────────── @@ -228,6 +234,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..c1157387 100644 --- a/frontend/web/src/features/partners/api/partnersApi.ts +++ b/frontend/web/src/features/partners/api/partnersApi.ts @@ -54,6 +54,10 @@ class HttpPartnersRepository implements IPartnersRepository { } // ── Platform Trading Partners ────────────── + deleteCertificateSecret(vaultRef: string): Promise { + return this.request(`/api/v1/platform/trading-partners/as2/certificates/secret?vault_ref=${encodeURIComponent(vaultRef)}`, { method: 'DELETE' }); + } + getPlatformPartners(): Promise { return this.request('/api/v1/platform/trading-partners/as2/trading-partners'); } @@ -124,6 +128,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..ba03bacd 100644 --- a/frontend/web/src/features/partners/components/As2PartnerDetails.tsx +++ b/frontend/web/src/features/partners/components/As2PartnerDetails.tsx @@ -1,20 +1,22 @@ -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'; +import { extractCertificateMaterial } from '../utils/certificate'; 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(); @@ -71,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, @@ -161,7 +152,7 @@ export function As2PartnerDetails({ partner, onCancel }: { partner: AS2Partner, control={control} render={({ field }) => ( - Status + Status Description + Issued + Expires + @@ -280,6 +274,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 { + return null; + } + }, [publicPem]); + const handleCopy = (text: string, label: string) => { navigator.clipboard.writeText(text); toast({ title: 'Copied', description: `${label} copied to clipboard.` }); @@ -327,21 +334,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/As2PartnersTable.tsx b/frontend/web/src/features/partners/components/As2PartnersTable.tsx index c65831a9..1b9319ce 100644 --- a/frontend/web/src/features/partners/components/As2PartnersTable.tsx +++ b/frontend/web/src/features/partners/components/As2PartnersTable.tsx @@ -9,13 +9,13 @@ import { } from '@tanstack/react-table'; import type { AS2Partner } from '../types'; -import { useDeletePlatformPartner, useUpdatePlatformPartnerMutation } from '../api/partnerHooks'; +import { useDeletePlatformPartnerMutation, useUpdatePlatformPartnerMutation } from '../api/partnerHooks'; import { Server, CheckCircle2 } from 'lucide-react'; import { As2PartnerDetails } from './As2PartnerDetails'; import { SharedRowActions } from './SharedRowActions'; function As2PartnerRowActions({ partner }: { partner: AS2Partner }) { - const deletePlatform = useDeletePlatformPartner(); + const deletePlatform = useDeletePlatformPartnerMutation(); const updatePlatform = useUpdatePlatformPartnerMutation(); const isDeleting = deletePlatform.isPending; diff --git a/frontend/web/src/features/partners/components/CreatePartnerModal.tsx b/frontend/web/src/features/partners/components/CreatePartnerModal.tsx index 5dd8b84c..543456cf 100644 --- a/frontend/web/src/features/partners/components/CreatePartnerModal.tsx +++ b/frontend/web/src/features/partners/components/CreatePartnerModal.tsx @@ -1,48 +1,71 @@ -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'; import { CertificateInput } from './CertificateInput'; -import { useCreatePlatformPartnerMutation } from '../api/partnerHooks'; -import { usePlatformConfig } from '@/features/platform/api/configHooks'; +import { useCreatePlatformPartnerMutation, useGenerateCertificateMutation, useDeleteCertificateSecretMutation } 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'; +import { extractCertificateMaterial } from '../utils/certificate'; 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 [generatedForAs2Id, setGeneratedForAs2Id] = useState(null); 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: platformConfig } = usePlatformConfig(); + const { data: platformSettings } = usePlatformSettings(); const { toast } = useToast(); const createPartner = useCreatePlatformPartnerMutation(); + const generateCert = useGenerateCertificateMutation(); + const deleteCertSecret = useDeleteCertificateSecretMutation(); - useEffect(() => { - if (isLocal && !url && platformConfig?.available_as2_receive_urls?.length) { - setUrl(platformConfig.available_as2_receive_urls[0]); + const handleCleanup = async () => { + if (privateKeyVaultRef) { + await deleteCertSecret.mutateAsync(privateKeyVaultRef); } - }, [isLocal, platformConfig, url]); + }; - const reset = () => { - setIsLocal(false); - setCertPem(''); - setAs2Id(''); - setUrl(''); + const reset = async () => { + // Only cleanup if we are abandoning an unsaved draft + 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) => { setIsOpen(open); - if (!open) reset(); + if (!open) { + reset(); + } }; const handleSubmit = async (e: React.FormEvent) => { e.preventDefault(); const data = new FormData(e.currentTarget); + const submittedAs2Id = data.get('as2_id') as string; if (!url || url.trim() === '') { toast({ title: 'Error', description: 'Receiving URL is required.', variant: 'destructive' }); @@ -56,19 +79,58 @@ export function CreatePartnerModal({ existingAs2Ids = [] }: { existingAs2Ids?: s return; } + // 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 + 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; + toast({ title: 'Warning', description: 'AS2 ID changed. Please regenerate the certificate.', variant: 'destructive' }); + return; + } + createPartner.mutate( { name: data.get('name') as string, type: 'AS2', - as2_id: data.get('as2_id') as string, + as2_id: submittedAs2Id, is_local: isLocal, url: url, - public_cert_pem: isLocal ? undefined : certPem, + // 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: () => { setIsOpen(false); - reset(); + // Don't call reset() here because we don't want to delete the saved secret + setIsLocal(false); + setCertPem(''); + setPrivateKeyVaultRef(null); + setGeneratedForAs2Id(null); + setAs2Id(''); + setUrl(''); }, }, ); @@ -93,9 +155,37 @@ export function CreatePartnerModal({ existingAs2Ids = [] }: { existingAs2Ids?: s type="button" role="switch" aria-checked={isLocal} - onClick={() => { - setIsLocal(!isLocal); - if (!isLocal) setCertPem(''); + onClick={async () => { + const nextIsLocal = !isLocal; + if (nextIsLocal) { + 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 { + 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(''); + } + } }} 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 +229,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; + } + if (privateKeyVaultRef) { + 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) => { + 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); + 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 }: + + + {seconds.map(s => ( + {s.label} + ))} + + +
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ ); +} diff --git a/frontend/web/src/features/platform/components/SchedulerDashboard.tsx b/frontend/web/src/features/platform/components/SchedulerDashboard.tsx new file mode 100644 index 00000000..ef030d8e --- /dev/null +++ b/frontend/web/src/features/platform/components/SchedulerDashboard.tsx @@ -0,0 +1,238 @@ +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 { 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, 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 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 [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' }; + }; + + 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 * * * *'); + } + setIntervalError(null); + }; + + const handleSaveSchedule = async () => { + if (!editingJob) return; + + try { + if (scheduleType === 'interval') { + 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; + 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 } }); + } + setEditingJob(null); + setIntervalError(null); + } catch (e: any) { + setIntervalError(e.message || 'Failed to update job schedule'); + } + }; + + if (jobsLoading) return
Loading Scheduler Dashboard...
; + + return ( +
+

Scheduler

+ + + +
+ Scheduled Jobs +

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

+
+
+ +
+
+ + + + + Name + Status + Schedule + Next Run At + Actions + + + + {jobs.map((job: any) => ( + + {job.name} + + + {job.status} + + + + {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'} + +
+ + +
+
+
+ ))} + {jobs.length === 0 && ( + + + No background jobs found. + + + )} +
+
+
+
+ + !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/frontend/web/src/routeTree.gen.ts b/frontend/web/src/routeTree.gen.ts index e7817ce6..d3dbeb8b 100644 --- a/frontend/web/src/routeTree.gen.ts +++ b/frontend/web/src/routeTree.gen.ts @@ -24,6 +24,7 @@ import { Route as TenantDevelopersRouteImport } from './routes/tenant/developers import { Route as TenantDashboardRouteImport } from './routes/tenant/dashboard' import { Route as PlatformUsersRouteImport } from './routes/platform/users' import { Route as PlatformTenantsRouteImport } from './routes/platform/tenants' +import { Route as PlatformSchedulerRouteImport } from './routes/platform/scheduler' import { Route as PlatformPartnershipsRouteImport } from './routes/platform/partnerships' import { Route as PlatformPartnersRouteImport } from './routes/platform/partners' import { Route as PlatformDashboardRouteImport } from './routes/platform/dashboard' @@ -104,6 +105,11 @@ const PlatformTenantsRoute = PlatformTenantsRouteImport.update({ path: '/tenants', getParentRoute: () => PlatformRoute, } as any) +const PlatformSchedulerRoute = PlatformSchedulerRouteImport.update({ + id: '/scheduler', + path: '/scheduler', + getParentRoute: () => PlatformRoute, +} as any) const PlatformPartnershipsRoute = PlatformPartnershipsRouteImport.update({ id: '/partnerships', path: '/partnerships', @@ -137,6 +143,7 @@ export interface FileRoutesByFullPath { '/platform/dashboard': typeof PlatformDashboardRoute '/platform/partners': typeof PlatformPartnersRoute '/platform/partnerships': typeof PlatformPartnershipsRoute + '/platform/scheduler': typeof PlatformSchedulerRoute '/platform/tenants': typeof PlatformTenantsRoute '/platform/users': typeof PlatformUsersRoute '/tenant/dashboard': typeof TenantDashboardRoute @@ -156,6 +163,7 @@ export interface FileRoutesByTo { '/platform/dashboard': typeof PlatformDashboardRoute '/platform/partners': typeof PlatformPartnersRoute '/platform/partnerships': typeof PlatformPartnershipsRoute + '/platform/scheduler': typeof PlatformSchedulerRoute '/platform/tenants': typeof PlatformTenantsRoute '/platform/users': typeof PlatformUsersRoute '/tenant/dashboard': typeof TenantDashboardRoute @@ -179,6 +187,7 @@ export interface FileRoutesById { '/platform/dashboard': typeof PlatformDashboardRoute '/platform/partners': typeof PlatformPartnersRoute '/platform/partnerships': typeof PlatformPartnershipsRoute + '/platform/scheduler': typeof PlatformSchedulerRoute '/platform/tenants': typeof PlatformTenantsRoute '/platform/users': typeof PlatformUsersRoute '/tenant/dashboard': typeof TenantDashboardRoute @@ -203,6 +212,7 @@ export interface FileRouteTypes { | '/platform/dashboard' | '/platform/partners' | '/platform/partnerships' + | '/platform/scheduler' | '/platform/tenants' | '/platform/users' | '/tenant/dashboard' @@ -222,6 +232,7 @@ export interface FileRouteTypes { | '/platform/dashboard' | '/platform/partners' | '/platform/partnerships' + | '/platform/scheduler' | '/platform/tenants' | '/platform/users' | '/tenant/dashboard' @@ -244,6 +255,7 @@ export interface FileRouteTypes { | '/platform/dashboard' | '/platform/partners' | '/platform/partnerships' + | '/platform/scheduler' | '/platform/tenants' | '/platform/users' | '/tenant/dashboard' @@ -373,6 +385,13 @@ declare module '@tanstack/react-router' { preLoaderRoute: typeof PlatformTenantsRouteImport parentRoute: typeof PlatformRoute } + '/platform/scheduler': { + id: '/platform/scheduler' + path: '/scheduler' + fullPath: '/platform/scheduler' + preLoaderRoute: typeof PlatformSchedulerRouteImport + parentRoute: typeof PlatformRoute + } '/platform/partnerships': { id: '/platform/partnerships' path: '/partnerships' @@ -427,6 +446,7 @@ interface PlatformRouteChildren { PlatformDashboardRoute: typeof PlatformDashboardRoute PlatformPartnersRoute: typeof PlatformPartnersRoute PlatformPartnershipsRoute: typeof PlatformPartnershipsRoute + PlatformSchedulerRoute: typeof PlatformSchedulerRoute PlatformTenantsRoute: typeof PlatformTenantsRoute PlatformUsersRoute: typeof PlatformUsersRoute } @@ -435,6 +455,7 @@ const PlatformRouteChildren: PlatformRouteChildren = { PlatformDashboardRoute: PlatformDashboardRoute, PlatformPartnersRoute: PlatformPartnersRoute, PlatformPartnershipsRoute: PlatformPartnershipsRoute, + PlatformSchedulerRoute: PlatformSchedulerRoute, PlatformTenantsRoute: PlatformTenantsRoute, PlatformUsersRoute: PlatformUsersRoute, } diff --git a/frontend/web/src/routes/platform.tsx b/frontend/web/src/routes/platform.tsx index 5ead4884..4886fd8f 100644 --- a/frontend/web/src/routes/platform.tsx +++ b/frontend/web/src/routes/platform.tsx @@ -9,7 +9,8 @@ import { Network, ChevronRight, Server, - LogOut + LogOut, + Clock, } from 'lucide-react' import { useDashboardData } from '@/features/dashboard/api/useDashboardData' @@ -99,6 +100,7 @@ export function AppLayout() {
System Admin
+ 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/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/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 09952c1c..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 @@ -339,12 +339,53 @@ def upgrade() -> None: unique=True, postgresql_where=sa.text("active = true"), ) + op.create_table( + "platform_settings", + sa.Column("key", sa.String(), nullable=False), + sa.Column("value", sa.JSON(), nullable=True), + sa.Column("created_at", sa.DateTime(timezone=True), nullable=False), + sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False), + sa.PrimaryKeyConstraint("key"), + ) + op.create_table( + "scheduled_jobs", + sa.Column("id", sa.Uuid(), nullable=False), + sa.Column("name", sa.String(), nullable=False), + sa.Column("payload", sa.JSON(), nullable=False), + sa.Column("status", sa.String(), nullable=False), + 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), + sa.Column("created_at", sa.DateTime(timezone=True), nullable=False), + sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False), + sa.PrimaryKeyConstraint("id"), + ) + op.create_index(op.f("ix_scheduled_jobs_name"), "scheduled_jobs", ["name"], unique=False) + op.create_index( + op.f("ix_scheduled_jobs_next_run_at"), "scheduled_jobs", ["next_run_at"], unique=False + ) + op.create_index(op.f("ix_scheduled_jobs_status"), "scheduled_jobs", ["status"], unique=False) # ### end Alembic commands ### 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/migrations/tenant/versions/f966e8446341_tenant_initial_schema.py b/libs/database/src/database/migrations/tenant/versions/f966e8446341_tenant_initial_schema.py index c5e16f69..a59ec56c 100644 --- a/libs/database/src/database/migrations/tenant/versions/f966e8446341_tenant_initial_schema.py +++ b/libs/database/src/database/migrations/tenant/versions/f966e8446341_tenant_initial_schema.py @@ -384,7 +384,7 @@ def upgrade() -> None: sa.Column("id", sa.UUID(), server_default=sa.text("gen_random_uuid()"), nullable=False), sa.Column("trace_id", sa.UUID(), nullable=False), sa.Column("direction", sa.String(length=50), nullable=False), - sa.Column("outbound_route_id", sa.UUID(), nullable=True), + sa.Column("trading_partner_id", sa.String(length=255), nullable=True), sa.Column("transaction_type", sa.String(length=50), nullable=True), sa.Column("standard", sa.String(length=50), nullable=True), sa.Column("sender_id", sa.String(length=255), nullable=True), @@ -406,10 +406,6 @@ def upgrade() -> None: sa.CheckConstraint( "(payload IS NOT NULL OR storage_uri IS NOT NULL)", name="chk_edi_json_data_or_uri" ), - sa.ForeignKeyConstraint( - ["outbound_route_id"], - ["outbound_routes.id"], - ), sa.PrimaryKeyConstraint("id"), ) op.create_index( @@ -420,7 +416,7 @@ def upgrade() -> None: postgresql_using="gin", ) op.create_index( - op.f("ix_edi_json_outbound_route_id"), "edi_json", ["outbound_route_id"], unique=False + op.f("ix_edi_json_trading_partner_id"), "edi_json", ["trading_partner_id"], unique=False ) op.create_index( "ix_edi_json_sender_recv", @@ -455,7 +451,7 @@ def upgrade() -> None: sa.Column("status_message", sa.Text(), nullable=True), sa.Column("state", sa.String(length=255), nullable=True), sa.Column("msg_headers", sa.Text(), nullable=True), - sa.Column("outbound_route_id", sa.UUID(), nullable=True), + sa.Column("trading_partner_id", sa.String(length=255), nullable=True), sa.Column("as2_sender_id", sa.String(length=255), nullable=True), sa.Column("as2_receiver_id", sa.String(length=255), nullable=True), sa.Column("interchange_control_no", sa.String(length=255), nullable=True), @@ -476,16 +472,12 @@ def upgrade() -> None: sa.CheckConstraint( "(edi_data IS NOT NULL OR storage_uri IS NOT NULL)", name="chk_edi_msg_data_or_uri" ), - sa.ForeignKeyConstraint( - ["outbound_route_id"], - ["outbound_routes.id"], - ), sa.PrimaryKeyConstraint("id"), ) op.create_index( - op.f("ix_edi_messages_outbound_route_id"), + op.f("ix_edi_messages_trading_partner_id"), "edi_messages", - ["outbound_route_id"], + ["trading_partner_id"], unique=False, ) op.create_index(op.f("ix_edi_messages_tenant_id"), "edi_messages", ["tenant_id"], unique=False) @@ -505,13 +497,13 @@ def downgrade() -> None: op.drop_index("ix_edi_msgs_sender_recv", table_name="edi_messages") op.drop_index(op.f("ix_edi_messages_trace_id"), table_name="edi_messages") op.drop_index(op.f("ix_edi_messages_tenant_id"), table_name="edi_messages") - op.drop_index(op.f("ix_edi_messages_outbound_route_id"), table_name="edi_messages") + op.drop_index(op.f("ix_edi_messages_trading_partner_id"), table_name="edi_messages") op.drop_table("edi_messages") op.drop_index(op.f("ix_edi_json_transaction_type"), table_name="edi_json") op.drop_index(op.f("ix_edi_json_trace_id"), table_name="edi_json") op.drop_index(op.f("ix_edi_json_tenant_id"), table_name="edi_json") op.drop_index("ix_edi_json_sender_recv", table_name="edi_json") - op.drop_index(op.f("ix_edi_json_outbound_route_id"), table_name="edi_json") + op.drop_index(op.f("ix_edi_json_trading_partner_id"), table_name="edi_json") op.drop_index("ix_edi_json_business_metadata", table_name="edi_json", postgresql_using="gin") op.drop_table("edi_json") op.drop_index( diff --git a/libs/database/src/database/models/__init__.py b/libs/database/src/database/models/__init__.py index cb5e41b8..c62e0ddf 100644 --- a/libs/database/src/database/models/__init__.py +++ b/libs/database/src/database/models/__init__.py @@ -1,6 +1,7 @@ from .control_plane import ( AS2Partner, AS2Partnership, + ControlPlaneOutbox, DatabaseShard, GlobalBase, SystemAuditLog, @@ -14,13 +15,11 @@ from .control_plane import ( OutboundRoute as GlobalOutboundRoute, ) -from .control_plane import ( - Outbox as GlobalOutbox, -) from .data_plane import ( AckReceipt, ApiGateway, AuditLog, + DataPlaneOutbox, EdiMessage, InboundRoute, Job, @@ -32,9 +31,8 @@ TenantBase, Webhook, ) -from .data_plane import ( - Outbox as TenantOutbox, -) +from .platform_settings import PlatformSettings +from .scheduled_job import ScheduledJob __all__ = [ # Global @@ -45,7 +43,7 @@ "TenantUser", "AS2Partner", "AS2Partnership", - "GlobalOutbox", + "ControlPlaneOutbox", "SystemAuditLog", "GlobalOutboundEdiHeader", "GlobalOutboundRoute", @@ -60,8 +58,11 @@ "EdiMessage", "ApiGateway", "Job", - "TenantOutbox", + "DataPlaneOutbox", "ProcessedEvent", "AuditLog", "AckReceipt", + # Scheduler + "ScheduledJob", + "PlatformSettings", ] diff --git a/libs/database/src/database/models/control_plane.py b/libs/database/src/database/models/control_plane.py index a872f40e..3ba42cff 100644 --- a/libs/database/src/database/models/control_plane.py +++ b/libs/database/src/database/models/control_plane.py @@ -148,7 +148,7 @@ class AS2Partnership(GlobalBase, AS2PartnershipMixin, TimestampMixin): ) -class Outbox(GlobalBase, OutboxMixin): +class ControlPlaneOutbox(GlobalBase, OutboxMixin): __tablename__ = "outbox" id: Mapped[PyUUID] = mapped_column( diff --git a/libs/database/src/database/models/data_plane.py b/libs/database/src/database/models/data_plane.py index 749d494f..c16e6404 100644 --- a/libs/database/src/database/models/data_plane.py +++ b/libs/database/src/database/models/data_plane.py @@ -202,9 +202,7 @@ class EdiMessage(TenantBase, TenantAwareMixin, TimestampMixin): status_message: Mapped[str | None] = mapped_column(Text, nullable=True) state: Mapped[str | None] = mapped_column(String(255), nullable=True) msg_headers: Mapped[str | None] = mapped_column(Text, nullable=True) - outbound_route_id: Mapped[PyUUID | None] = mapped_column( - UUID(as_uuid=True), ForeignKey("outbound_routes.id"), nullable=True, index=True - ) + trading_partner_id: Mapped[str | None] = mapped_column(String(255), nullable=True, index=True) as2_sender_id: Mapped[str | None] = mapped_column(String(255), nullable=True) as2_receiver_id: Mapped[str | None] = mapped_column(String(255), nullable=True) @@ -236,9 +234,7 @@ class EdiJson(TenantBase, TenantAwareMixin, TimestampMixin): trace_id: Mapped[PyUUID] = mapped_column(UUID(as_uuid=True), nullable=False, index=True) direction: Mapped[str] = mapped_column(String(50), nullable=False) # INBOUND, OUTBOUND - outbound_route_id: Mapped[PyUUID | None] = mapped_column( - UUID(as_uuid=True), ForeignKey("outbound_routes.id"), nullable=True, index=True - ) + trading_partner_id: Mapped[str | None] = mapped_column(String(255), nullable=True, index=True) transaction_type: Mapped[str | None] = mapped_column(String(50), nullable=True, index=True) standard: Mapped[str | None] = mapped_column(String(50), nullable=True) sender_id: Mapped[str | None] = mapped_column(String(255), nullable=True) @@ -304,7 +300,7 @@ class Job(TenantBase, TenantAwareMixin, TimestampMixin): error_message: Mapped[str | None] = mapped_column(Text, nullable=True) -class Outbox(TenantBase, TenantAwareMixin, OutboxMixin): +class DataPlaneOutbox(TenantBase, TenantAwareMixin, OutboxMixin): __tablename__ = "outbox" id: Mapped[PyUUID] = mapped_column( 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..a0ce0957 --- /dev/null +++ b/libs/database/src/database/models/scheduled_job.py @@ -0,0 +1,45 @@ +import uuid +from datetime import UTC, 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") + + 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=lambda: datetime.now(UTC) + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), + default=lambda: datetime.now(UTC), + onupdate=lambda: datetime.now(UTC), + ) diff --git a/libs/domain/src/domain/models.py b/libs/domain/src/domain/models.py index b48ee673..69cab3f5 100644 --- a/libs/domain/src/domain/models.py +++ b/libs/domain/src/domain/models.py @@ -56,7 +56,7 @@ class EdiRecordBase(BaseModel): class EdiJsonDomainModel(EdiRecordBase): - outbound_route_id: UUID | None = None + trading_partner_id: str | None = None transaction_type: str | None = None standard: str | None = None sender_id: str | None = None @@ -77,7 +77,7 @@ class EdiMessageDomainModel(EdiRecordBase): gs_sender_id: str | None = None gs_receiver_id: str | None = None inbound_route_id: UUID | None = None - outbound_route_id: UUID | None = None + trading_partner_id: str | None = None edi_data: str | None = None # Populated from DB or S3 storage_uri: str | None = None 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/adapters/repository.py b/libs/pipeline/src/pipeline/adapters/repository.py index 32e12223..d4edb668 100644 --- a/libs/pipeline/src/pipeline/adapters/repository.py +++ b/libs/pipeline/src/pipeline/adapters/repository.py @@ -1,11 +1,11 @@ import uuid -from datetime import datetime +from datetime import UTC, datetime from typing import Any from config.settings import AppSettings from database.encryption import db_encryption from database.models import ApiGateway, EdiMessage -from database.models import TenantOutbox as Outbox +from database.models import DataPlaneOutbox as Outbox from database.models.data_plane import ( AS2Partner, AS2Partnership, @@ -65,7 +65,7 @@ async def update_edi_message_status(self, trace_id: str, status: str) -> None: .scalar_subquery() ) ) - .values(status=status, updated_at=datetime.utcnow()) + .values(status=status, updated_at=datetime.now(UTC)) ) await self.session.execute(stmt) @@ -109,7 +109,7 @@ async def save_edi_message( receiver_id: str | None = None, gs_sender_id: str | None = None, gs_receiver_id: str | None = None, - outbound_route_id: str | None = None, + trading_partner_id: str | None = None, tenant_id: int | None = None, ) -> None: storage_uri = None @@ -135,7 +135,7 @@ async def save_edi_message( "gs_receiver_id": gs_receiver_id, "storage_uri": storage_uri, "status": status, - "outbound_route_id": uuid.UUID(outbound_route_id) if outbound_route_id else None, + "trading_partner_id": trading_partner_id, } if tenant_id is not None: record_kwargs["tenant_id"] = tenant_id @@ -192,7 +192,7 @@ async def save_edi_json( record_kwargs = { "trace_id": uuid.UUID(trace_id), "direction": direction, - "outbound_route_id": uuid.UUID(partnership_id) if partnership_id else None, + "trading_partner_id": partnership_id, "transaction_type": transaction_type, "standard": standard, "sender_id": sender_id, 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 3b231af4..80534ad3 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 @@ -51,14 +57,17 @@ async def deliver(self, trace_id: str, partner_id: str, edi_msg: EdiMessageDomai 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( @@ -66,16 +75,18 @@ async def deliver(self, trace_id: str, partner_id: str, edi_msg: EdiMessageDomai 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/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 4b759059..c8cfb409 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. @@ -32,12 +32,22 @@ async def deliver(self, trace_id: str) -> None: direction = edi_msg.direction - if direction == "OUTBOUND" and edi_msg.outbound_route_id: - route = await self.repository.get_outbound_route(str(edi_msg.outbound_route_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, + ) if not route: - logger.error(f"Configured outbound route for {edi_msg.outbound_route_id} not found") + logger.error( + f"Configured outbound route for trading_partner_id={edi_msg.trading_partner_id} not found" + ) raise ValueError( - f"Configured outbound route for {edi_msg.outbound_route_id} not found" + f"Configured outbound route for trading_partner_id={edi_msg.trading_partner_id} not found" ) else: sender_id = edi_msg.sender_id @@ -56,11 +66,13 @@ async def deliver(self, trace_id: str) -> None: 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) + 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/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/core/saga.py b/libs/pipeline/src/pipeline/core/saga.py index 4419166b..873c8aac 100644 --- a/libs/pipeline/src/pipeline/core/saga.py +++ b/libs/pipeline/src/pipeline/core/saga.py @@ -47,12 +47,12 @@ async def handle_transform_completed(self, payload: dict[str, Any]) -> None: # For OUTBOUND, the input is EdiJson. Its terminal state is TRANSFORMED. import uuid - outbound_route_id = payload.get("outbound_route_id") + trading_partner_id = payload.get("trading_partner_id") # Pack update kwargs (omitting None values if not present) update_kwargs = {} - if outbound_route_id: - update_kwargs["outbound_route_id"] = uuid.UUID(outbound_route_id) + if trading_partner_id: + update_kwargs["trading_partner_id"] = trading_partner_id if "standard" in payload: update_kwargs["standard"] = payload["standard"] if "isa_sender_id" in payload: diff --git a/libs/pipeline/src/pipeline/core/transformation/inbound.py b/libs/pipeline/src/pipeline/core/transformation/inbound.py index 0a5b8441..db4e3889 100644 --- a/libs/pipeline/src/pipeline/core/transformation/inbound.py +++ b/libs/pipeline/src/pipeline/core/transformation/inbound.py @@ -110,7 +110,7 @@ async def transform(self, trace_id: str) -> None: await self.repository.save_edi_json( trace_id=trace_id, direction=MessageDirection.INBOUND, - partnership_id=partnership_id_str, + partnership_id=None, transaction_type=txn_type, standard=standard, sender_id=sender_id, diff --git a/libs/pipeline/src/pipeline/core/transformation/outbound.py b/libs/pipeline/src/pipeline/core/transformation/outbound.py index 259ac319..37d67570 100644 --- a/libs/pipeline/src/pipeline/core/transformation/outbound.py +++ b/libs/pipeline/src/pipeline/core/transformation/outbound.py @@ -32,33 +32,24 @@ async def transform(self, trace_id: str) -> None: if not edi_json: raise ValueError(f"No EdiJson record found for trace_id={trace_id}") - outbound_route_id = str(edi_json.outbound_route_id) if edi_json.outbound_route_id else None + trading_partner_id = edi_json.trading_partner_id tenant_id = edi_json.tenant_id business_metadata = edi_json.business_metadata or {} routing_meta = business_metadata.get("_routing") or {} - trading_partner_id = routing_meta.get("trading_partner_id") + # Prefer explicit trading_partner_id on the record; fall back to business_metadata routing hint + if not trading_partner_id: + trading_partner_id = routing_meta.get("trading_partner_id") - if outbound_route_id: - route_config = await self.repository.get_outbound_edi_header_by_route_or_partner( - route_id=outbound_route_id - ) - # Also get the route to link to EdiMessage - outbound_route = await self.repository.get_outbound_route(outbound_route_id) - else: - if not trading_partner_id or not tenant_id: - raise ValueError( - f"No routing info available (trading_partner_id/tenant_id) for trace_id={trace_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 ) - # Get the route to link to EdiMessage outbound_route = await self.repository.get_outbound_route_by_trading_partner_id( trading_partner_id=trading_partner_id, tenant_id=tenant_id ) - if outbound_route: - outbound_route_id = outbound_route.get("id") if not route_config or not outbound_route: raise ValueError( @@ -131,7 +122,7 @@ async def transform(self, trace_id: str) -> None: receiver_id=route_config.get("isa_receiver_id"), gs_sender_id=route_config.get("gs_sender_id"), gs_receiver_id=route_config.get("gs_receiver_id"), - outbound_route_id=outbound_route_id, + trading_partner_id=trading_partner_id, tenant_id=edi_json.tenant_id, ) @@ -144,7 +135,7 @@ async def transform(self, trace_id: str) -> None: payload={ "trace_id": trace_id, "direction": MessageDirection.OUTBOUND, - "outbound_route_id": outbound_route_id, + "trading_partner_id": trading_partner_id, "standard": standard, "isa_sender_id": route_config.get("isa_sender_id"), "isa_receiver_id": route_config.get("isa_receiver_id"), 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/src/pipeline/ports/repository.py b/libs/pipeline/src/pipeline/ports/repository.py index 5f9a721c..f20dbf4d 100644 --- a/libs/pipeline/src/pipeline/ports/repository.py +++ b/libs/pipeline/src/pipeline/ports/repository.py @@ -26,7 +26,7 @@ async def save_edi_message( receiver_id: str | None = None, gs_sender_id: str | None = None, gs_receiver_id: str | None = None, - outbound_route_id: str | None = None, + trading_partner_id: str | None = None, tenant_id: int | None = None, ) -> None: """Stores a raw EDI message.""" diff --git a/libs/pipeline/tests/fakes.py b/libs/pipeline/tests/fakes.py index e299c21f..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) @@ -233,9 +250,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/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/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..eca9ef56 --- /dev/null +++ b/libs/scheduler/pyproject.toml @@ -0,0 +1,24 @@ +[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", + "croniter>=6.2.4", +] + +[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..e985b860 --- /dev/null +++ b/libs/scheduler/src/scheduler/adapters/repository.py @@ -0,0 +1,179 @@ +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, 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( + id=record.id, + name=record.name, + 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, + locked_by=record.locked_by, + created_at=record.created_at, + updated_at=record.updated_at, + ) + + async def claim_next_jobs(self, worker_id: str, limit: int) -> list[Job]: + """ + Uses SKIP LOCKED to safely claim the next PENDING or ready jobs. + """ + 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.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( + ScheduledJob.next_run_at.asc().nulls_first(), ScheduledJob.created_at.asc() + ) + .limit(limit) + .with_for_update(skip_locked=True) + ) + + result = await session.execute(stmt) + records = result.scalars().all() + + if not records: + return [] + + 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 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 = ( + 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, + interval_seconds: int | 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, + interval_seconds=interval_seconds, + ) + 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) + + 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 new file mode 100644 index 00000000..3a3ce1c7 --- /dev/null +++ b/libs/scheduler/src/scheduler/core/service.py @@ -0,0 +1,127 @@ +import asyncio +import contextlib +import logging +from typing import Any + +from scheduler.ports.publisher import MessagePublisherPort +from scheduler.ports.repository import JobRepositoryPort + +logger = logging.getLogger(__name__) + + +class SchedulerWorkerService: + 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.max_concurrent_jobs = max_concurrent_jobs + self._is_running = False + self._task: asyncio.Task[None] | None = None + 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} with concurrency {self.max_concurrent_jobs}" + ) + 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 + + 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: + 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()} + + 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: + 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..8244ff74 --- /dev/null +++ b/libs/scheduler/src/scheduler/domain/models.py @@ -0,0 +1,61 @@ +import uuid +from dataclasses import dataclass, field +from datetime import datetime +from enum import StrEnum +from typing import Any, cast + + +class JobStatus(StrEnum): + PENDING = "PENDING" + RUNNING = "RUNNING" + COMPLETED = "COMPLETED" + FAILED = "FAILED" + PAUSED = "PAUSED" + + +class JobName(StrEnum): + OUTBOX_SWEEPER = "outbox_sweeper" + DATA_RETENTION_CLEANUP = "data_retention_cleanup" + + +@dataclass +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 new file mode 100644 index 00000000..1da86ef6 --- /dev/null +++ b/libs/scheduler/src/scheduler/ports/handler.py @@ -0,0 +1,10 @@ +import abc + +from scheduler.domain.models import Job + + +class JobHandlerPort(abc.ABC): + @abc.abstractmethod + 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 new file mode 100644 index 00000000..b69ea9a4 --- /dev/null +++ b/libs/scheduler/src/scheduler/ports/repository.py @@ -0,0 +1,42 @@ +import abc +import uuid +from datetime import datetime, timedelta +from typing import Any + +from scheduler.domain.models import Job + + +class JobRepositoryPort(abc.ABC): + @abc.abstractmethod + 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 + + @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, + 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/api_token_repository.py b/services/api/src/api/adapters/api_token_repository.py index cecc81ff..d8923377 100644 --- a/services/api/src/api/adapters/api_token_repository.py +++ b/services/api/src/api/adapters/api_token_repository.py @@ -115,7 +115,7 @@ async def get_tenant_id_by_credentials(self, client_id: str, secret_hash: str) - from sqlalchemy import update - now = datetime.now(UTC).replace(tzinfo=None) + now = datetime.now(UTC) result = await self.session.execute( # type: ignore select(ApiToken).where( ApiToken.client_id == client_id, diff --git a/services/api/src/api/adapters/http/dtos.py b/services/api/src/api/adapters/http/dtos.py index 961d780c..d9a605a9 100644 --- a/services/api/src/api/adapters/http/dtos.py +++ b/services/api/src/api/adapters/http/dtos.py @@ -2,7 +2,7 @@ from typing import Annotated, Any, Literal from uuid import UUID -from pydantic import BaseModel, Field, HttpUrl, field_validator, model_validator +from pydantic import BaseModel, ConfigDict, Field, HttpUrl, field_validator, model_validator # --------------------------------------------------------------------------- # Partner Creation Requests @@ -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): @@ -306,6 +309,17 @@ def model_post_init(self, __context: Any) -> None: object.__setattr__(self, "id", self.partner_id) +class GenerateCertRequest(BaseModel): + as2_id: str = Field(..., max_length=255, description="AS2 ID to use as Common Name") + + +class GenerateCertResponse(BaseModel): + public_cert_pem: str = Field(..., description="Public certificate in PEM format") + private_key_vault_ref: str = Field( + ..., description="Vault reference for the generated private key" + ) + + class AS2TradingPartnerResponse(BaseModel): id: str name: str @@ -354,6 +368,8 @@ class RouteResponse(BaseModel): class BaseRouteItem(BaseModel): + model_config = ConfigDict(from_attributes=True) + route_id: UUID trading_partner_id: str | None = None name: str diff --git a/services/api/src/api/adapters/outbox_repository.py b/services/api/src/api/adapters/outbox_repository.py index ea5f0987..a8ede093 100644 --- a/services/api/src/api/adapters/outbox_repository.py +++ b/services/api/src/api/adapters/outbox_repository.py @@ -3,15 +3,13 @@ from uuid import UUID from api.ports.outbox_repository import OutboxRepositoryPort -from database.base_repository import BaseSqlAlchemyRepository -from database.models.control_plane import Outbox as GlobalOutbox -from sqlalchemy.ext.asyncio import AsyncSession +from database.base_repository import GlobalSqlAlchemyRepository, TenantSqlAlchemyRepository +from database.models.control_plane import ControlPlaneOutbox -class SqlAlchemyOutboxRepository(OutboxRepositoryPort, BaseSqlAlchemyRepository): - def __init__(self, session: AsyncSession, model_class: Any = GlobalOutbox) -> None: - self.session = session - self.model_class = model_class +class SqlAlchemyOutboxRepositoryMixin: + session: Any + model_class: Any async def publish_outbox_event( self, @@ -32,3 +30,33 @@ async def publish_outbox_event( self.session.add(record) await self.session.flush() return event_id + + +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 + by the CDC Sweeper, which is configurable through the Scheduler UI. + """ + + def __init__(self, session: Any) -> None: + from database.models.data_plane import DataPlaneOutbox + + super().__init__(session) + self.model_class = DataPlaneOutbox 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 7374fe05..f9c09030 100644 --- a/services/api/src/api/adapters/transaction_repository.py +++ b/services/api/src/api/adapters/transaction_repository.py @@ -28,10 +28,10 @@ async def create_edi_message(self, tenant_id: int, payload: dict[str, Any]) -> U async def publish_outbox_event( self, tenant_id: int, event_type: str, payload: dict[str, Any], idempotency_key: UUID ) -> UUID: - from database.models.data_plane import Outbox + from database.models.data_plane import DataPlaneOutbox event_id = uuid.uuid4() - record = Outbox( + record = DataPlaneOutbox( id=event_id, tenant_id=tenant_id, idempotency_key=idempotency_key, @@ -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,23 +177,37 @@ 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 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): @@ -276,7 +298,7 @@ async def get_transaction(self, tenant_id: int, trace_id: UUID) -> TransactionDe encryption_algorithm=edi_msg.encryption_algorithm, compression=getattr(edi_msg, "compression", None), inbound_route_id=getattr(edi_msg, "inbound_route_id", None), - outbound_route_id=getattr(edi_msg, "outbound_route_id", None), + trading_partner_id=getattr(edi_msg, "trading_partner_id", None), status=getattr(edi_msg, "status", "RECEIVED"), edi_data=getattr(edi_msg, "edi_data", None), interchange_control_no=getattr(edi_msg, "interchange_control_no", None), @@ -296,6 +318,7 @@ async def get_transaction(self, tenant_id: int, trace_id: UUID) -> TransactionDe id=j.id, trace_id=j.trace_id, status=j.status, + trading_partner_id=getattr(j, "trading_partner_id", None), error_message=getattr(j, "error_message", None), interchange_control_number=getattr(j, "interchange_control_number", None), group_control_number=getattr(j, "group_control_number", None), diff --git a/services/api/src/api/cdc_relay.py b/services/api/src/api/cdc_relay.py index fe6cec96..fa406d64 100644 --- a/services/api/src/api/cdc_relay.py +++ b/services/api/src/api/cdc_relay.py @@ -20,10 +20,9 @@ _TRANSFORM_QUEUE_EVENT_TYPES: frozenset[str] = frozenset( { PipelineEventType.TRANSFORM_EVENT, + PipelineEventType.COMPUTE_TRANSFORM_EVENT, PipelineEventType.TRANSFORM_COMPLETED, PipelineEventType.DELIVERY_COMPLETED, - "json.received", - "edi_message.received", } ) diff --git a/services/api/src/api/core/services/as2_partner_service.py b/services/api/src/api/core/services/as2_partner_service.py index cbfaa9aa..05fdcb35 100644 --- a/services/api/src/api/core/services/as2_partner_service.py +++ b/services/api/src/api/core/services/as2_partner_service.py @@ -28,7 +28,7 @@ async def create_as2_partner( logger.info(f"Provisioning AS2 partner {cmd.name} for tenant {tenant_id}") partner_id = await self.uow.as2_partners.create_as2_identity(tenant_id=tenant_id, cmd=cmd) - await self.uow.outbox.publish_outbox_event( + await self.uow.control_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=ProvisioningEventType.AS2_PARTNER_CREATED, payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, @@ -52,7 +52,7 @@ async def update_as2_partner( if not updated_partner: raise ValueError("Partner not found after update") - await self.uow.outbox.publish_outbox_event( + await self.uow.control_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=ProvisioningEventType.AS2_PARTNER_UPDATED, payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, @@ -69,7 +69,7 @@ async def update_as2_partner( async def delete_as2_partner(self, tenant_id: int, partner_id: UUID) -> None: logger.info(f"Deleting AS2 partner {partner_id} for tenant {tenant_id}") await self.uow.as2_partners.delete_as2_identity(tenant_id, partner_id) - await self.uow.outbox.publish_outbox_event( + await self.uow.control_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=ProvisioningEventType.AS2_PARTNER_DELETED, payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, @@ -91,7 +91,7 @@ async def rotate_certificates( if not updated_partner: raise ValueError("Partner not found after certificate rotation") - await self.uow.outbox.publish_outbox_event( + await self.uow.control_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=ProvisioningEventType.AS2_PARTNER_UPDATED, payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, diff --git a/services/api/src/api/core/services/as2_partnership_service.py b/services/api/src/api/core/services/as2_partnership_service.py index cbd8f2bd..efb5d7ee 100644 --- a/services/api/src/api/core/services/as2_partnership_service.py +++ b/services/api/src/api/core/services/as2_partnership_service.py @@ -41,7 +41,7 @@ async def create_as2_partnership( partner_id = await self.uow.as2_partnerships.create_as2_partnership( tenant_id=tenant_id, cmd=cmd ) - await self.uow.outbox.publish_outbox_event( + await self.uow.control_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=ProvisioningEventType.AS2_PARTNERSHIP_CREATED, payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, @@ -77,7 +77,7 @@ async def update_as2_partnership( await self.uow.as2_partnerships.update_as2_partnership( tenant_id=tenant_id, partnership_id=partnership_id, cmd=cmd ) - await self.uow.outbox.publish_outbox_event( + await self.uow.control_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=ProvisioningEventType.AS2_PARTNERSHIP_UPDATED, payload={"partner_id": str(partnership_id), "tenant_id": tenant_id}, @@ -97,7 +97,7 @@ async def update_as2_partnership( async def delete_as2_partnership(self, tenant_id: int, partnership_id: UUID) -> None: logger.info(f"Deleting AS2 partnership {partnership_id} for tenant {tenant_id}") await self.uow.as2_partnerships.delete_as2_partnership(tenant_id, partnership_id) - await self.uow.outbox.publish_outbox_event( + await self.uow.control_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=ProvisioningEventType.AS2_PARTNERSHIP_DELETED, payload={"partner_id": str(partnership_id), "tenant_id": tenant_id}, diff --git a/services/api/src/api/core/services/inbound_route_service.py b/services/api/src/api/core/services/inbound_route_service.py index 4befa1ec..484a8bfc 100644 --- a/services/api/src/api/core/services/inbound_route_service.py +++ b/services/api/src/api/core/services/inbound_route_service.py @@ -26,7 +26,7 @@ def __init__(self, uow: UnitOfWork) -> None: async def create_inbound_route(self, tenant_id: int, cmd: CreateInboundRouteCmd) -> RouteEntity: logger.info(f"Creating Inbound Route for sender {cmd.isa_sender_id} in tenant {tenant_id}") route_id = await self.uow.inbound_routes.create_inbound_route(tenant_id=tenant_id, cmd=cmd) - await self.uow.outbox.publish_outbox_event( + await self.uow.control_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=ProvisioningEventType.INBOUND_ROUTE_CREATED, payload={"route_id": str(route_id), "tenant_id": tenant_id}, @@ -38,7 +38,7 @@ async def update_inbound_route( ) -> bool: res = await self.uow.inbound_routes.update_inbound_route(tenant_id, route_id, cmd) if res: - await self.uow.outbox.publish_outbox_event( + await self.uow.control_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=ProvisioningEventType.INBOUND_ROUTE_UPDATED, payload={"route_id": str(route_id), "tenant_id": tenant_id}, @@ -48,7 +48,7 @@ async def update_inbound_route( async def delete_inbound_route(self, tenant_id: int, route_id: UUID) -> bool: res = await self.uow.inbound_routes.delete_inbound_route(tenant_id, route_id) if res: - await self.uow.outbox.publish_outbox_event( + await self.uow.control_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=ProvisioningEventType.INBOUND_ROUTE_DELETED, payload={"route_id": str(route_id), "tenant_id": tenant_id}, diff --git a/services/api/src/api/core/services/outbound_route_service.py b/services/api/src/api/core/services/outbound_route_service.py index d6883c9d..66aa1fef 100644 --- a/services/api/src/api/core/services/outbound_route_service.py +++ b/services/api/src/api/core/services/outbound_route_service.py @@ -32,7 +32,7 @@ async def create_outbound_route( route_id = await self.uow.outbound_routes.create_outbound_route( tenant_id=tenant_id, cmd=cmd ) - await self.uow.outbox.publish_outbox_event( + await self.uow.control_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=ProvisioningEventType.OUTBOUND_ROUTE_CREATED, payload={"route_id": str(route_id), "tenant_id": tenant_id}, @@ -44,7 +44,7 @@ async def update_outbound_route( ) -> bool: res = await self.uow.outbound_routes.update_outbound_route(tenant_id, route_id, cmd) if res: - await self.uow.outbox.publish_outbox_event( + await self.uow.control_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=ProvisioningEventType.OUTBOUND_ROUTE_UPDATED, payload={"route_id": str(route_id), "tenant_id": tenant_id}, @@ -54,7 +54,7 @@ async def update_outbound_route( async def delete_outbound_route(self, tenant_id: int, route_id: UUID) -> bool: res = await self.uow.outbound_routes.delete_outbound_route(tenant_id, route_id) if res: - await self.uow.outbox.publish_outbox_event( + await self.uow.control_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=ProvisioningEventType.OUTBOUND_ROUTE_DELETED, payload={"route_id": str(route_id), "tenant_id": tenant_id}, diff --git a/services/api/src/api/core/services/routing_resolver.py b/services/api/src/api/core/services/routing_resolver.py index 300a5ba4..e08853c1 100644 --- a/services/api/src/api/core/services/routing_resolver.py +++ b/services/api/src/api/core/services/routing_resolver.py @@ -27,7 +27,7 @@ async def resolve_routing_context( self, msg: Any, edi_jsons: list[Any] ) -> tuple[str | None, str | None]: if ( - getattr(msg, "outbound_route_id", None) + getattr(msg, "trading_partner_id", None) or getattr(msg, "direction", None) == Direction.OUTBOUND ): return await self._resolve_outbound_routing(msg, edi_jsons) @@ -40,12 +40,15 @@ async def _resolve_outbound_routing( Resolves outbound routing by first checking explicit route overrides, then falling back to business_metadata from the EDI JSON. """ - # 1. Try to resolve via outbound_route_id - if getattr(msg, "outbound_route_id", None) and self.tenant_session: + # 1. Try to resolve via trading_partner_id on the message + if getattr(msg, "trading_partner_id", None) and self.tenant_session: try: route = ( await self.tenant_session.execute( - select(OutboundRoute).where(OutboundRoute.id == msg.outbound_route_id) + select(OutboundRoute).where( + OutboundRoute.trading_partner_id == msg.trading_partner_id, + OutboundRoute.active.is_(True), + ) ) ).scalar_one_or_none() diff --git a/services/api/src/api/core/services/sftp_partner_service.py b/services/api/src/api/core/services/sftp_partner_service.py index ad96ebfb..de17ea2c 100644 --- a/services/api/src/api/core/services/sftp_partner_service.py +++ b/services/api/src/api/core/services/sftp_partner_service.py @@ -27,7 +27,7 @@ def __init__(self, uow: UnitOfWork) -> None: async def create_sftp_partner(self, tenant_id: int, cmd: CreateSFTPPartnerCmd) -> PartnerEntity: logger.info(f"Creating SFTP partner {cmd.name} for tenant {tenant_id}") partner_id = await self.uow.sftp_partners.create_sftp_partner(tenant_id=tenant_id, cmd=cmd) - await self.uow.outbox.publish_outbox_event( + await self.uow.control_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=ProvisioningEventType.SFTP_PARTNER_CREATED, payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, @@ -70,7 +70,7 @@ async def update_sftp_partner( ) update_hash = hashlib.sha256(str(cmd).encode()).hexdigest() - await self.uow.outbox.publish_outbox_event( + await self.uow.control_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=ProvisioningEventType.SFTP_PARTNER_UPDATED, payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, diff --git a/services/api/src/api/core/services/webhook_service.py b/services/api/src/api/core/services/webhook_service.py index 9e94c5fa..89c09f79 100644 --- a/services/api/src/api/core/services/webhook_service.py +++ b/services/api/src/api/core/services/webhook_service.py @@ -31,7 +31,7 @@ def __init__(self, uow: UnitOfWork) -> None: async def create_webhook(self, tenant_id: int, cmd: CreateWebhookCmd) -> PartnerEntity: logger.info("Webhook creating", extra={"tenant_id": tenant_id, "webhook_name": cmd.name}) partner_id = await self.uow.webhooks.create_webhook(tenant_id=tenant_id, cmd=cmd) - await self.uow.outbox.publish_outbox_event( + await self.uow.control_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=ProvisioningEventType.WEBHOOK_CREATED, payload={"partner_id": str(partner_id), "tenant_id": tenant_id}, @@ -57,7 +57,7 @@ async def update_webhook( ) result = await self.uow.webhooks.update_webhook(tenant_id, webhook_id, name, active, url) if result: - await self.uow.outbox.publish_outbox_event( + await self.uow.control_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=ProvisioningEventType.WEBHOOK_UPDATED, payload={"partner_id": str(webhook_id), "tenant_id": tenant_id}, @@ -70,7 +70,7 @@ async def delete_webhook(self, tenant_id: int, webhook_id: UUID) -> bool: ) result = await self.uow.webhooks.delete_webhook(tenant_id, webhook_id) if result: - await self.uow.outbox.publish_outbox_event( + await self.uow.control_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=ProvisioningEventType.WEBHOOK_DELETED, payload={"partner_id": str(webhook_id), "tenant_id": tenant_id}, diff --git a/services/api/src/api/core/uow.py b/services/api/src/api/core/uow.py index 41281c5d..d026b3f4 100644 --- a/services/api/src/api/core/uow.py +++ b/services/api/src/api/core/uow.py @@ -6,7 +6,11 @@ from api.adapters.edi_header_repository import SqlAlchemyEdiHeaderRepository from api.adapters.inbound_route_repository import SqlAlchemyInboundRouteRepository from api.adapters.outbound_route_repository import SqlAlchemyOutboundRouteRepository -from api.adapters.outbox_repository import SqlAlchemyOutboxRepository +from api.adapters.outbox_repository import ( + SqlAlchemyControlPlaneOutboxRepository, + SqlAlchemyDataPlaneOutboxRepository, +) +from api.adapters.platform_settings_repository import SqlAlchemyPlatformSettingsRepository from api.adapters.sftp_repository import SqlAlchemySFTPPartnerRepository from api.adapters.tenant_repository import SqlAlchemyTenantRepository from api.adapters.transaction_repository import SqlAlchemyTransactionRepository @@ -40,11 +44,13 @@ def __init__( self.as2_partnerships = SqlAlchemyAS2PartnershipRepository(gs) self.inbound_routes = SqlAlchemyInboundRouteRepository(gs) self.outbound_routes = SqlAlchemyOutboundRouteRepository(gs) - self.outbox = SqlAlchemyOutboxRepository(gs) + self.control_plane_outbox = SqlAlchemyControlPlaneOutboxRepository(gs) + self._data_plane_outbox = SqlAlchemyDataPlaneOutboxRepository(ts) if ts else None self.sftp_partners = SqlAlchemySFTPPartnerRepository(gs) self.tenants = SqlAlchemyTenantRepository(gs) self.webhooks = SqlAlchemyWebhookRepository(gs) self.edi_headers = SqlAlchemyEdiHeaderRepository(gs) + self.platform_settings = SqlAlchemyPlatformSettingsRepository(gs) self._transactions = SqlAlchemyTransactionRepository(ts) if ts else None @@ -54,6 +60,12 @@ def transactions(self) -> SqlAlchemyTransactionRepository: raise RuntimeError("Transaction repository requires an active tenant session.") return self._transactions + @property + def data_plane_outbox(self) -> SqlAlchemyDataPlaneOutboxRepository: + if not self._data_plane_outbox: + raise RuntimeError("Tenant outbox repository requires an active tenant session.") + return self._data_plane_outbox + async def __aenter__(self) -> Self: return self diff --git a/services/api/src/api/domain/models.py b/services/api/src/api/domain/models.py index b6168f53..c084c000 100644 --- a/services/api/src/api/domain/models.py +++ b/services/api/src/api/domain/models.py @@ -322,7 +322,7 @@ class EdiMessageDTO: encryption_algorithm: str | None = None compression: str | None = None inbound_route_id: UUID | None = None - outbound_route_id: UUID | None = None + trading_partner_id: str | None = None status: str = "RECEIVED" edi_data: str | None = None interchange_control_no: str | None = None @@ -343,6 +343,7 @@ class EdiJsonDTO: id: UUID trace_id: UUID status: str + trading_partner_id: str | None = None error_message: str | None = None interchange_control_number: str | None = None group_control_number: str | None = None diff --git a/services/api/src/api/main.py b/services/api/src/api/main.py index a8a25d74..4322a1fb 100644 --- a/services/api/src/api/main.py +++ b/services/api/src/api/main.py @@ -29,6 +29,9 @@ transactions, webhooks, ) +from api.routers import ( + platform as platform_admin, +) from api.routers.developers import api_tokens from api.routers.trading_partners import as2_receive, platform @@ -97,6 +100,7 @@ async def validation_exception_handler( app.include_router(trading_partners.router) app.include_router(webhooks.router) app.include_router(platform.router) +app.include_router(platform_admin.router) app.include_router(routes.router) app.include_router(edi_headers.router) app.include_router(edi_tools.router) 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..e19da731 --- /dev/null +++ b/services/api/src/api/routers/platform/__init__.py @@ -0,0 +1,10 @@ +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", 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 new file mode 100644 index 00000000..b0f050c6 --- /dev/null +++ b/services/api/src/api/routers/platform/scheduler.py @@ -0,0 +1,265 @@ +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 + 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 + updated_at: datetime + + 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 + + +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.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) + + 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." + ) + 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: + 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, + 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=job_def.max_retries, + 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.""" + 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.""" + async with uow: + await uow.platform_settings.set_config(key, request.value) + await uow.commit() + return ConfigResponse(key=key, value=request.value) diff --git a/services/api/src/api/routers/trading_partners/platform/__init__.py b/services/api/src/api/routers/trading_partners/platform/__init__.py index 55d20a76..b1adc565 100644 --- a/services/api/src/api/routers/trading_partners/platform/__init__.py +++ b/services/api/src/api/routers/trading_partners/platform/__init__.py @@ -8,7 +8,7 @@ from fastapi import APIRouter, Depends from api.dependencies import require_platform_admin -from api.routers.trading_partners.platform import as2_partners, as2_partnerships, config +from api.routers.trading_partners.platform import as2_partners, as2_partnerships, settings _PREFIX = "/api/v1/platform/trading-partners" @@ -20,4 +20,4 @@ router.include_router(as2_partners.router) router.include_router(as2_partnerships.router) -router.include_router(config.router) +router.include_router(settings.router) 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 914dc511..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 @@ -7,6 +7,8 @@ from api.adapters.http.dtos import ( AS2TradingPartnerResponse, CreateAS2TradingPartnerRequest, + GenerateCertRequest, + GenerateCertResponse, UpdateAS2TradingPartnerRequest, ) from api.core.services import AS2PartnerService @@ -25,6 +27,61 @@ router = APIRouter(tags=["Platform Partners - AS2"]) +@router.post( + "/as2/certificates/generate", + response_model=GenerateCertResponse, + status_code=status.HTTP_200_OK, +) +async def generate_certificate( + request: GenerateCertRequest, + vault: VaultPort = Depends(get_vault), +) -> Any: + """ + Generates a new self-signed AS2 certificate and stores the private key in Vault. + Returns the public cert PEM and the vault reference for the private key. + """ + 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.as2_id.replace(" ", "_").lower(), + ) + + return GenerateCertResponse( + public_cert_pem=public_cert_bytes.decode("utf-8"), + private_key_vault_ref=private_key_vault_ref, + ) + + +@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), + 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) + + @router.post( "/as2/trading-partners", response_model=AS2TradingPartnerResponse, @@ -39,24 +96,35 @@ 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 + commit_success = False + try: async with uow: - public_cert_pem = request.public_cert_pem - private_key_vault_ref = request.private_key_vault_ref - if request.is_local: - # Auto-generate self-signed cert - 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 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, @@ -73,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") @@ -85,10 +155,14 @@ async def create_platform_as2_partner( url=p.url, active=p.active, ) - except IntegrityError as e: - if request.is_local and private_key_vault_ref: + except Exception as e: + if auto_generated and private_key_vault_ref and not commit_success: 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/routers/trading_partners/platform/config.py b/services/api/src/api/routers/trading_partners/platform/settings.py similarity index 82% rename from services/api/src/api/routers/trading_partners/platform/config.py rename to services/api/src/api/routers/trading_partners/platform/settings.py index a0f7e2a0..b64e1f0c 100644 --- a/services/api/src/api/routers/trading_partners/platform/config.py +++ b/services/api/src/api/routers/trading_partners/platform/settings.py @@ -4,7 +4,7 @@ from fastapi import APIRouter from pydantic import BaseModel -router = APIRouter(tags=["Platform Config"]) +router = APIRouter(tags=["Platform Settings"]) class SupportedAlgorithm(BaseModel): @@ -12,20 +12,20 @@ class SupportedAlgorithm(BaseModel): label: str -class PlatformConfigResponse(BaseModel): +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=PlatformConfigResponse) -async def get_platform_config() -> Any: +@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 PlatformConfigResponse( + return PlatformSettingsResponse( available_as2_receive_urls=[f"{base_url}/api/v1/as2/receive"], supported_as2_encryption_algorithms=[ SupportedAlgorithm(value="AES256", label="AES-256-CBC"), diff --git a/services/api/src/api/services/api_receiver_service.py b/services/api/src/api/services/api_receiver_service.py index 7eb58096..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() @@ -88,7 +102,7 @@ async def process_api_edi_json( # 5. Drop Outbox event for Worker to transform from domain.events import PipelineEventType - await self.uow.outbox.publish_outbox_event( + await self.uow.data_plane_outbox.publish_outbox_event( tenant_id=tenant_id, event_type=PipelineEventType.TRANSFORM_EVENT, payload={ diff --git a/services/api/src/api/services/as2_receiver_service.py b/services/api/src/api/services/as2_receiver_service.py index ec5fbdff..f52cf065 100644 --- a/services/api/src/api/services/as2_receiver_service.py +++ b/services/api/src/api/services/as2_receiver_service.py @@ -11,6 +11,7 @@ from as2_core.message import AS2Message from as2_core.parser import parse_as2_request from database.models.control_plane import DatabaseShard, Tenant +from domain.events import PipelineEventType from security.smime import decrypt_payload, verify_signature from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession @@ -373,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 = { @@ -385,9 +388,9 @@ 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="edi_message.received", + event_type=PipelineEventType.TRANSFORM_EVENT, payload=outbox_payload, idempotency_key=msg_id, ) 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..1379837f 100644 --- a/services/api/tests/test_api_receiver_service.py +++ b/services/api/tests/test_api_receiver_service.py @@ -16,11 +16,68 @@ 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() + 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.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 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() + create_args, create_kwargs = mock_uow.transactions.create_edi_json.await_args + assert create_kwargs["payload"]["transaction_type"] == "850" + + +@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() + 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) + 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/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/api/tests/test_scheduler.py b/services/api/tests/test_scheduler.py new file mode 100644 index 00000000..7fec8d4e --- /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": "outbox_sweeper", "interval_seconds": 60, "payload": {}}, + ) + assert resp.status_code == 200 + assert mock_uow.global_session.add.called + assert resp.json()["name"] == "outbox_sweeper" + + +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": "outbox_sweeper"}, + ) + 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..e2b32c23 100644 --- a/services/as2_server/scripts/seed.py +++ b/services/as2_server/scripts/seed.py @@ -97,6 +97,50 @@ 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.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/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/adapters/sqs_poller.py b/services/workers/orchestrator/src/worker/adapters/sqs_poller.py new file mode 100644 index 00000000..dac0c6cb --- /dev/null +++ b/services/workers/orchestrator/src/worker/adapters/sqs_poller.py @@ -0,0 +1,80 @@ +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 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..32e53091 --- /dev/null +++ b/services/workers/orchestrator/src/worker/adapters/sqs_publisher.py @@ -0,0 +1,111 @@ +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}") + raise + + 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}") + raise 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..90969ed2 --- /dev/null +++ b/services/workers/orchestrator/src/worker/core/security.py @@ -0,0 +1,106 @@ +import ipaddress +import logging +import socket +from contextlib import contextmanager +from contextvars import ContextVar +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 + + +_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 new file mode 100644 index 00000000..5418978f --- /dev/null +++ b/services/workers/orchestrator/src/worker/core/tenant_resolver.py @@ -0,0 +1,52 @@ +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, 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.monotonic() + self._sweep(now) + + 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) + 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 new file mode 100644 index 00000000..cbd9a186 --- /dev/null +++ b/services/workers/orchestrator/src/worker/data/handlers.py @@ -0,0 +1,230 @@ +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 ssrf_safe_context +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=ssrf_safe_context) + sftp_adapter = ParamikoSftpDeliveryAdapter() + vault_adapter = WorkerVaultAdapter() + as2_adapter = HttpxAS2DeliveryAdapter(validator=ssrf_safe_context) + + 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 + + # 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="DELIVERING") + ) + await session.commit() + + # 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, + ) + + 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 + + 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 137029d5..16d75f2a 100644 --- a/services/workers/orchestrator/src/worker/data/main.py +++ b/services/workers/orchestrator/src/worker/data/main.py @@ -1,368 +1,146 @@ 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 sqlalchemy import select -from worker.adapters.vault import WorkerVaultAdapter +from scheduler.adapters.repository import SqlAlchemyJobRepository +from scheduler.core.service import SchedulerWorkerService +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 main() -> None: + settings = get_settings() + aws_endpoint = os.getenv("AWS_ENDPOINT_URL") + s3_bucket = "soopaedi-dev" -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) + db_router = DatabaseRouter(global_db_url=settings.database.global_url) + resolver = TenantResolver(db_router) - 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, + 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"), ) - 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, + transform_task = asyncio.create_task( + poll_sqs_queue( + MessageQueueName.TRANSFORM_ORCHESTRATION_QUEUE, + pipeline_processor, + aws_endpoint, ) - 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, + 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"), ) - # 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"] + deliver_task = asyncio.create_task( + poll_sqs_queue( + MessageQueueName.DELIVER_QUEUE, + delivery_processor, + aws_endpoint, + ) + ) - logger.info(f"Started polling {queue_name} ({queue_url})") + # Start Scheduler worker + from sqlalchemy.ext.asyncio import create_async_engine - while True: - response = await sqs.receive_message( - QueueUrl=queue_url, - MaxNumberOfMessages=10, - WaitTimeSeconds=20, - ) + 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, + ) - 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") + scheduler_service = SchedulerWorkerService( + scheduler_repo, publisher=message_publisher, worker_id=f"orchestrator-{os.getpid()}" + ) - 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 + 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 - 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) + registry = JobHandlerRegistry() + registry.register( + JobName.OUTBOX_SWEEPER.value, DataPlaneOutboxSweeperJobHandler(db_router, message_publisher) + ) + registry.register( + JobName.DATA_RETENTION_CLEANUP.value, DataRetentionCleanupJobHandler(db_router) + ) - # 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}" - ) + import functools - 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) + from worker.data.scheduled_jobs_handler import process_scheduled_job + scheduled_jobs_processor = functools.partial(process_scheduled_job, registry=registry) -async def main() -> None: - settings = get_settings() - aws_endpoint = os.getenv("AWS_ENDPOINT_URL") - s3_bucket = "soopaedi-dev" - - db_router = DatabaseRouter(global_db_url=settings.database.global_url) - resolver = TenantResolver(db_router) - - transform_task = asyncio.create_task( + scheduled_jobs_task = asyncio.create_task( poll_sqs_queue( - MessageQueueName.TRANSFORM_ORCHESTRATION_QUEUE, - process_pipeline_event, - resolver, - db_router, - s3_bucket, - aws_endpoint, - ) - ) - deliver_task = asyncio.create_task( - poll_sqs_queue( - MessageQueueName.DELIVER_QUEUE, - process_delivery, - resolver, - db_router, - s3_bucket, + "edi-orchestrator-jobs", + scheduled_jobs_processor, aws_endpoint, ) ) - await asyncio.gather(transform_task, deliver_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 new file mode 100644 index 00000000..4083cb1c --- /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}") + 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) + + 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..165c4929 --- /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}") + raise + + 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 new file mode 100644 index 00000000..0d831679 --- /dev/null +++ b/services/workers/orchestrator/src/worker/jobs/outbox_sweeper.py @@ -0,0 +1,143 @@ +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 +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 +from worker.ports.message_publisher import MessagePublisherPort + +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, +} + +_BATCH_SIZE = 100 +_CONCURRENCY_LIMIT = 5 + + +class DataPlaneOutboxSweeperJobHandler(JobHandlerPort): + def __init__(self, db_router: DatabaseRouter, message_publisher: MessagePublisherPort) -> None: + self.db_router = db_router + self.message_publisher = message_publisher + + 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) + + 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: + 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) + + logger.info( + f"[DataPlaneOutboxSweeper] Sweep complete. Total events forwarded: {total_processed}" + ) + + 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) -> 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) + ) + 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(): + 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, + }, + } + ) + + 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. + # 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/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/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..acb2465e 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,76 @@ 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=SHARD_1_URL) + 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 == SHARD_1_URL + + # 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 +179,49 @@ 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 + + +@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..98333e16 --- /dev/null +++ b/services/workers/orchestrator/tests/test_security.py @@ -0,0 +1,59 @@ +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("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"): + import socket + + 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(): + with pytest.raises(ValueError), ssrf_safe_context("ftp://example.com"): + pass 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/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 diff --git a/uv.lock b/uv.lock index 8a052e85..7ea3496c 100644 --- a/uv.lock +++ b/uv.lock @@ -25,6 +25,7 @@ members = [ "orchestrator-worker", "patches", "pipeline", + "scheduler", "security", "transformer", ] @@ -896,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" @@ -2134,6 +2147,7 @@ dependencies = [ { name = "database" }, { name = "domain" }, { name = "pipeline" }, + { name = "scheduler" }, { name = "sqlalchemy" }, ] @@ -2145,6 +2159,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 +2929,25 @@ 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 = "croniter" }, + { name = "database" }, + { name = "pydantic" }, + { name = "sqlalchemy" }, +] + +[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" }, +] + [[package]] name = "security" version = "0.1.0"