diff --git a/.gitignore b/.gitignore index eca270c9..47932757 100644 --- a/.gitignore +++ b/.gitignore @@ -11,6 +11,10 @@ __pycache__/ *.egg-info/ dist/ .env +config/tenants.yaml +config/tenants.yaml.tmp +config/playground_tenants.yaml +config/playground_tenants.yaml.tmp .DS_Store **/.DS_Store .coverage @@ -46,3 +50,7 @@ launch.json # log mcp_debug.log .cursor/ + +# nw-mcp-builder +nw-mcp-builder/fixtures/ +nw-mcp-builder/out/ diff --git a/docs/configuration.md b/docs/configuration.md index 11033b8d..d8d07ed0 100644 --- a/docs/configuration.md +++ b/docs/configuration.md @@ -65,6 +65,17 @@ copy sample.env .env | `NW_JWT_ISSUER` | Expected JWT `iss` claim when any `*_JWT_SECRET` is set | _(required with JWT secret)_ | | `NW_SMTP_ALLOWED_HOSTS` | Optional comma-separated SMTP relay hostnames permitted for `smtp.send_email` (recommended for production) | _(unset = env relay only)_ | +### Multi-tenancy + +| Variable | Description | Default | +|----------|-------------|---------| +| `NW_MULTITENANCY_ENABLED` | When `true`, resolve tenant from header / `NW_TENANT_ID` / JWT and require a tenant (missing → error). When `false`, always `__default__`. | `false` | +| `NW_TENANT_ID` | **MCP stdio only** — pins the process to one tenant. Do not set on multi-tenant streamable-http (use `X-Tenant-ID` instead). Required when multitenancy is enabled for stdio. | _(unset)_ | +| `NW_TENANT_ID_HEADER` | HTTP/gRPC header name for tenant id (case-insensitive) | `X-Tenant-ID` | +| `NW_TENANTS_PATH` | Path to the YAML file that persists runtime named configs + tenant secret overlays (`config/tenants.yaml` by default; gitignored) | `config/tenants.yaml` | + +Named-tenant secrets use `NW_{TENANT}_{CONNECTOR}_{CONFIG}_{KEY}` (one credential vault per named config). MCP transport details: [mcp-servers.md](mcp-servers.md#multi-tenancy-mcp). + --- ## Configuration File (`config/connectors.yaml`) diff --git a/playground/README.md b/playground/README.md index 6cbbf49d..6f7e1758 100644 --- a/playground/README.md +++ b/playground/README.md @@ -94,21 +94,49 @@ The Agentic Workflow panel displays the active transport as a pill: Set the mode before starting the REST API: ```powershell -# Buffered stdio mode +# Playground agent: in-process MCP (recommended for local MT testing) +$env:MODE="API" $env:NW_MCP_TRANSPORT="stdio" +# Leave TOOLHIVE_MCP_URL empty unless ToolHive / a separate MCP server is running uv run node-wire ``` ```powershell -# Streamable HTTP mode +# Streamable HTTP UI + remote MCP proxy (requires something listening on the URL) +$env:MODE="API" $env:NW_MCP_TRANSPORT="streamable-http" -$env:NW_MCP_HOST="127.0.0.1" -$env:NW_MCP_PORT="8081" -$env:NW_MCP_PATH="/mcp" +$env:TOOLHIVE_MCP_URL="http://127.0.0.1:8081/mcp" +# In another terminal: start MCP — see docs/mcp.md +uv run node-wire +``` + +After changing `NW_MCP_TRANSPORT` / `TOOLHIVE_MCP_URL`, restart the backend and hard refresh the browser so the latest `app.js` and transport status are loaded. + +**Modes (do not confuse):** + +| Goal | What to run | +|------|-------------| +| Playground + connector scenarios + Agent (tenant-aware) | `MODE=API` → `http://127.0.0.1:8000/playground/` | +| Standalone MCP tools (Inspector / Claude) | `python -m agents.mcp_entrypoint` (`NW_MCP_TRANSPORT=stdio` or `streamable-http` on `:8081`) — see [docs/mcp.md](../docs/mcp.md) | +| Full binding as MCP process | `MODE=MCP` with same transport vars | + +If `TOOLHIVE_MCP_URL` points at `:8081` but nothing is listening, Agent chat fails with `All connection attempts failed`. Clear the URL or start an MCP server; by default the playground falls back to **in-process** MCP when the proxy cannot list tools. + +#### Multitenancy (`NW_MULTITENANCY_ENABLED`) + +Defaults to off (legacy single-tenant). When enabled: + +```powershell +$env:NW_MULTITENANCY_ENABLED="true" uv run node-wire ``` -After changing `NW_MCP_TRANSPORT`, restart the backend and hard refresh the browser so the latest `app.js` and transport status are loaded. +- **Tenant ID required**: connector/scenario calls without `X-Tenant-ID` return **400**. Explicit `__default__` is allowed. +- **Header**: Tenant dropdown (existing tenants) and Config dropdown appear when multitenancy is on. Header **Add config** is hidden; use **Add config** on each System Connector page. +- **Config dropdown**: lists configs for the **active connector** only under the selected tenant. Switching connectors clears a name that does not exist for the new connector. +- **Agentic Workflow**: sends the same `X-Tenant-ID` (and optional `config_name` query). Local agent MCP runs **in-process** against the playground factory. Streamable-HTTP / ToolHive proxy URLs receive `X-Tenant-ID` on each MCP HTTP request. (Agent Add-config UI is unchanged / out of scope for this flow.) +- **Per-connector Add config**: On a connector page, **Add config** opens a modal for that connector only (tenant free-text for new tenants, config name, default flag, and varying credentials). Each named config has its **own** credential vault (`NW_{TENANT}_{CONNECTOR}_{CONFIG}_{KEY}`). Shared host env (e.g. `EPIC_FHIR_BASE_URL`, `EPIC_TOKEN_URL`) is copied into that config’s secret overlay when omitted. Persist file: gitignored `config/tenants.yaml` (holds secrets — do not commit). +- **Tenant in logs**: When multitenancy is on, server INFO lines include `tenant_id` / `config_name` for Agent chat, connector scenarios, config mutations, REST connector calls, and MCP tool resolution. The playground Technical Audit panel also prints Tenant/Config for those actions. #### Testing the MCP server with Inspector diff --git a/playground/app.js b/playground/app.js index b2c49525..12ea5a19 100644 --- a/playground/app.js +++ b/playground/app.js @@ -329,6 +329,7 @@ document.addEventListener('DOMContentLoaded', () => { Back to Workspace `; } + syncMultitenancyUi(); log('Switched to AI Agent mode (MCP + LLM)', 'system'); } else if (view === 'connector-apps-menu') { rootSelectionView.classList.add('hidden'); @@ -345,6 +346,7 @@ document.addEventListener('DOMContentLoaded', () => { Back to Workspace `; } + syncMultitenancyUi(); log('Opened Connector Apps menu', 'system'); } else if (view === 'ext-patient-viewer') { document.getElementById('connector-apps-selection-view').classList.add('hidden'); @@ -365,6 +367,7 @@ document.addEventListener('DOMContentLoaded', () => { `; } setMode('ext-patient-viewer'); + syncMultitenancyUi(); } else { rootSelectionView.classList.add('hidden'); layoutMain.classList.remove('hidden'); @@ -384,6 +387,7 @@ document.addEventListener('DOMContentLoaded', () => { Back to Workspace `; } + syncMultitenancyUi(); log('Switched to Connectors view', 'system'); } }); @@ -618,6 +622,10 @@ document.addEventListener('DOMContentLoaded', () => { if (slackPanel) slackPanel.classList.add('hidden'); if (extViewerPanel) extViewerPanel.classList.add('hidden'); + if (typeof syncMultitenancyUi === 'function') { + syncMultitenancyUi(); + } + if (mode === 'ehr') { ehrPanel.classList.remove('hidden'); @@ -724,6 +732,8 @@ document.addEventListener('DOMContentLoaded', () => { if (backSelectionBtn) backSelectionBtn.classList.add('hidden'); if (backToConnectorsBtn) backToConnectorsBtn.classList.remove('hidden'); setMode(mode); + refreshConfigDropdown(mode); + syncMultitenancyUi(); }); }); @@ -738,6 +748,7 @@ document.addEventListener('DOMContentLoaded', () => { connectorStatus.textContent = 'Connectors Ready'; tagline.textContent = 'Enterprise Integration Suite'; document.documentElement.style.setProperty('--brand-accent', '#2563eb'); + syncMultitenancyUi(); log('Returned to Connectors list', 'system'); } }); @@ -779,6 +790,582 @@ document.addEventListener('DOMContentLoaded', () => { gdriveActionSelect.addEventListener('change', scheduleGdriveDriveActionSyncFromUser); } + // --- Multitenancy feature flag (fetched once on page load) --- + let _multitenancyEnabled = false; + + function syncMultitenancyUi() { + const headerTenancy = document.getElementById('header-tenancy'); + if (headerTenancy) { + headerTenancy.classList.toggle('hidden', !_multitenancyEnabled); + } + const addConfigBtn = document.getElementById('add-config-btn'); + if (addConfigBtn) { + // Simplified: always hide header Add config (Agent work later). + addConfigBtn.classList.add('hidden'); + addConfigBtn.hidden = true; + } + const connectorAddBtn = document.getElementById('connector-add-config-btn'); + if (connectorAddBtn) { + const onConnector = + !!_modeToConnectorId[currentMode] && + playgroundView && + !playgroundView.classList.contains('hidden'); + connectorAddBtn.classList.toggle('hidden', !_multitenancyEnabled || !onConnector); + } + if (!_multitenancyEnabled) { + closeConfigAdminModal(); + } + } + + async function openConfigAdminModal() { + const modal = document.getElementById('config-admin-modal'); + if (!modal) return; + const connectorId = currentConnectorId(); + if (!connectorId) { + log('Open a connector page before adding config.', 'error'); + return; + } + const tip = document.getElementById('config-admin-connector-tip'); + if (tip) tip.textContent = connectorId; + + const tenantId = getPlaygroundTenantId(); + const configName = getPlaygroundConfigName(); + if (_cfgTenantInput) { + _cfgTenantInput.value = tenantId || ''; + } + + modal.classList.remove('hidden'); + setConfigAdminError(''); + await refreshModalConfigSelect(tenantId, connectorId, configName); + await applyModalConfigSelection(connectorId); + refreshConfigDropdown(currentMode); + } + + function closeConfigAdminModal() { + const modal = document.getElementById('config-admin-modal'); + if (modal) modal.classList.add('hidden'); + setConfigAdminError(''); + } + + function syncCfgNameGroupVisibility() { + const select = document.getElementById('cfg-select'); + const nameGroup = document.getElementById('cfg-name-group'); + const nameEl = document.getElementById('cfg-name'); + if (!select || !nameGroup || !nameEl) return; + const isNew = !select.value; + nameGroup.classList.toggle('hidden', !isNew); + if (!isNew) { + nameEl.value = select.value; + nameEl.readOnly = true; + } else { + nameEl.readOnly = false; + if (!nameEl.value.trim()) nameEl.value = 'test'; + } + } + + async function listConfigsForTenantConnector(tenantId, connectorId) { + if (!tenantId || !connectorId) return []; + const headers = { 'Content-Type': 'application/json', 'X-Tenant-ID': tenantId }; + const res = await fetch( + `/v1/connectors/${encodeURIComponent(connectorId)}/configs`, + { headers } + ); + if (!res.ok) return []; + const data = await res.json(); + return Array.isArray(data.configs) + ? data.configs + : Array.isArray(data) + ? data + : []; + } + + async function refreshModalConfigSelect(tenantId, connectorId, preferredName) { + const select = document.getElementById('cfg-select'); + if (!select) return; + const prev = preferredName != null ? preferredName : select.value; + select.innerHTML = ''; + if (!tenantId || !connectorId) { + syncCfgNameGroupVisibility(); + return; + } + try { + const configs = await listConfigsForTenantConnector(tenantId, connectorId); + configs.forEach((cfg) => { + const opt = document.createElement('option'); + opt.value = cfg.name || ''; + opt.textContent = (cfg.name || '') + (cfg.default ? ' (default)' : ''); + select.appendChild(opt); + }); + if (prev && [...select.options].some((o) => o.value === prev)) { + select.value = prev; + } else if (prev && configs.length === 0) { + // Deleted / unknown name: stay on new, keep typed name. + select.value = ''; + const nameEl = document.getElementById('cfg-name'); + if (nameEl && prev) nameEl.value = prev; + } else { + select.value = ''; + } + } catch (_) { + // ignore + } + syncCfgNameGroupVisibility(); + } + + async function applyModalConfigSelection(connectorId) { + const select = document.getElementById('cfg-select'); + const nameEl = document.getElementById('cfg-name'); + const defaultEl = document.getElementById('cfg-default'); + const tenantId = modalTenantId(); + const selected = select && select.value ? select.value : ''; + syncCfgNameGroupVisibility(); + + let existingKeys = []; + if (defaultEl) defaultEl.value = 'true'; + + if (tenantId && selected && connectorId) { + if (nameEl) nameEl.value = selected; + try { + const headers = { 'Content-Type': 'application/json', 'X-Tenant-ID': tenantId }; + const cfgRes = await fetch( + `/v1/connectors/${encodeURIComponent(connectorId)}/configs/${encodeURIComponent(selected)}`, + { headers } + ); + if (cfgRes.ok) { + const doc = await cfgRes.json(); + if (defaultEl && typeof doc.default === 'boolean') { + defaultEl.value = doc.default ? 'true' : 'false'; + } + const secRes = await fetch( + `/v1/connectors/${encodeURIComponent(connectorId)}/secrets?config_name=${encodeURIComponent(selected)}`, + { headers } + ); + if (secRes.ok) { + const sec = await secRes.json(); + existingKeys = Array.isArray(sec.keys) ? sec.keys : []; + } + } + } catch (_) { + // Prefill is best-effort. + } + } + renderSecretFields(connectorId || currentConnectorId(), existingKeys); + } + + function modalSelectedConfigName() { + const select = document.getElementById('cfg-select'); + if (select && select.value) return select.value.trim(); + return (document.getElementById('cfg-name')?.value || 'test').trim(); + } + + async function loadFeatureFlags() { + try { + const res = await fetch('/playground/feature-flags'); + if (res.ok) { + const flags = await res.json(); + _multitenancyEnabled = !!flags.multitenancy_enabled; + } + } catch (_) { + // If endpoint is unreachable, keep default (disabled). + } + syncMultitenancyUi(); + if (_multitenancyEnabled) { + await refreshTenantDropdown(); + await refreshConfigDropdown(currentMode); + } + } + + // Header Tenant is a select of existing tenants; modal Tenant ID is free-text for new ones. + const _tenantInput = document.getElementById('playground-tenant-id'); + const _cfgTenantInput = document.getElementById('cfg-tenant-id'); + let _tenantSyncLock = false; + + function syncTenantInputs(source) { + if (_tenantSyncLock) return; + _tenantSyncLock = true; + try { + const value = source && source.value != null ? String(source.value).trim() : ''; + if (_cfgTenantInput && source !== _cfgTenantInput) { + // Do not overwrite free-text while user types a new tenant in the modal. + if (source === _tenantInput) _cfgTenantInput.value = value; + } + if (_tenantInput && source === _cfgTenantInput && value) { + ensureTenantOption(value, true); + } + } finally { + _tenantSyncLock = false; + } + } + + function ensureTenantOption(tenantId, selectIt) { + if (!_tenantInput || !tenantId) return; + let found = false; + for (const opt of _tenantInput.options) { + if (opt.value === tenantId) { + found = true; + break; + } + } + if (!found) { + const opt = document.createElement('option'); + opt.value = tenantId; + opt.textContent = tenantId; + _tenantInput.appendChild(opt); + } + if (selectIt) _tenantInput.value = tenantId; + } + + async function refreshTenantDropdown() { + if (!_multitenancyEnabled || !_tenantInput) return; + const prev = _tenantInput.value; + try { + const res = await fetch('/v1/tenants'); + if (!res.ok) return; + const data = await res.json(); + const tenants = Array.isArray(data.tenants) ? data.tenants : []; + _tenantInput.innerHTML = ''; + tenants.forEach((tid) => { + const opt = document.createElement('option'); + opt.value = tid; + opt.textContent = tid; + _tenantInput.appendChild(opt); + }); + if (prev && tenants.includes(prev)) { + _tenantInput.value = prev; + } else if (prev) { + ensureTenantOption(prev, true); + } + } catch (_) { + // ignore + } + } + + let _tenantDebounce = null; + function onTenantChanged() { + clearTimeout(_tenantDebounce); + _tenantDebounce = setTimeout(() => refreshConfigDropdown(currentMode), 200); + if (_cfgTenantInput && _tenantInput) { + syncTenantInputs(_tenantInput); + } + } + if (_tenantInput) { + _tenantInput.addEventListener('change', onTenantChanged); + } + if (_cfgTenantInput) { + _cfgTenantInput.addEventListener('input', () => syncTenantInputs(_cfgTenantInput)); + let _cfgTenantDebounce = null; + _cfgTenantInput.addEventListener('change', () => { + clearTimeout(_cfgTenantDebounce); + _cfgTenantDebounce = setTimeout(async () => { + const connectorId = currentConnectorId(); + const tenantId = modalTenantId(); + await refreshModalConfigSelect(tenantId, connectorId, ''); + await applyModalConfigSelection(connectorId); + }, 150); + }); + _cfgTenantInput.addEventListener('blur', async () => { + const connectorId = currentConnectorId(); + const tenantId = modalTenantId(); + await refreshModalConfigSelect(tenantId, connectorId, document.getElementById('cfg-select')?.value || ''); + await applyModalConfigSelection(connectorId); + }); + } + + document.getElementById('cfg-select')?.addEventListener('change', async () => { + await applyModalConfigSelection(currentConnectorId()); + }); + + // Map playground mode names to connector_id strings used by the config API. + const _modeToConnectorId = { + ehr: 'fhir_epic', + cerner: 'fhir_cerner', + gdrive: 'google_drive', + itops: 'http_generic', + slack: 'slack', + stripe: 'stripe', + salesforce: 'salesforce', + }; + + function currentConnectorId() { + return _modeToConnectorId[currentMode] || ''; + } + + // Self-contained auth/config templates matching config/connectors.yaml (secret refs only). + const _connectorCreateDocs = { + http_generic: { + config: {}, + auth: {}, + }, + google_drive: { + config: {}, + auth: { + provider: 'service_account', + sa_json_secret: 'GOOGLE_DRIVE_SA_JSON', + scopes: ['https://www.googleapis.com/auth/drive'], + }, + }, + fhir_epic: { + config: { + base_url: 'https://fhir.epic.sandbox/api/FHIR/R4', + }, + auth: { + provider: 'oauth2', + grant_method: 'private_key_jwt', + token_url_secret: 'EPIC_TOKEN_URL', + client_id_secret: 'EPIC_CLIENT_ID', + private_key_secret: 'EPIC_PRIVATE_KEY', + kid_secret: 'EPIC_KID', + algorithm: 'RS384', + }, + }, + fhir_cerner: { + config: { + base_url: 'https://fhir-ehr-code.cerner.com/r4/your-tenant-id', + }, + auth: { + provider: 'oauth2', + grant_method: 'private_key_jwt', + token_url_secret: 'CERNER_TOKEN_URL', + client_id_secret: 'CERNER_CLIENT_ID', + private_key_secret: 'CERNER_PRIVATE_KEY', + kid_secret: 'CERNER_KID', + algorithm: 'RS384', + scopes_secret: 'CERNER_SCOPES', + scopes: [ + 'system/Patient.read', + 'system/Encounter.read', + 'system/DocumentReference.read', + 'system/DocumentReference.write', + ], + }, + }, + slack: { + config: {}, + auth: { + provider: 'static_token', + secret_key: 'SLACK_BOT_TOKEN', + }, + }, + stripe: { + config: {}, + auth: { + provider: 'static_token', + secret_key: 'stripe_api_key', + header_name: 'Authorization', + prefix: '', + }, + }, + salesforce: { + config: {}, + auth: { + provider: 'oauth2', + grant_method: 'refresh_token', + token_url_secret: 'SALESFORCE_TOKEN_URL', + client_id_secret: 'SALESFORCE_CLIENT_ID', + client_secret_secret: 'SALESFORCE_CLIENT_SECRET', + refresh_token_secret: 'SALESFORCE_REFRESH_TOKEN', + }, + }, + }; + + // Varying-only secret fields shown in Add config (shared env auto-copied server-side). + // Format / required secret validation is server-side only (tenant_store). + const _connectorSecretFields = { + google_drive: [ + { key: 'GOOGLE_DRIVE_SA_JSON', label: 'Service account JSON', type: 'textarea', required: true }, + ], + fhir_epic: [ + { key: 'EPIC_CLIENT_ID', label: 'Client ID', type: 'text', required: true }, + { key: 'EPIC_PRIVATE_KEY', label: 'Private key (PEM)', type: 'textarea', required: true }, + { key: 'EPIC_KID', label: 'Key ID (kid)', type: 'text', required: true }, + ], + fhir_cerner: [ + { key: 'CERNER_CLIENT_ID', label: 'Client ID', type: 'text', required: true }, + { key: 'CERNER_PRIVATE_KEY', label: 'Private key (PEM)', type: 'textarea', required: true }, + { key: 'CERNER_KID', label: 'Key ID (kid)', type: 'text', required: true }, + { key: 'CERNER_SCOPES', label: 'Scopes (space-separated)', type: 'text', required: false }, + ], + slack: [ + { key: 'SLACK_BOT_TOKEN', label: 'Bot token', type: 'text', required: true }, + ], + stripe: [ + { key: 'stripe_api_key', label: 'API key', type: 'text', required: true }, + ], + salesforce: [ + { key: 'SALESFORCE_CLIENT_ID', label: 'Client ID', type: 'text', required: true }, + { key: 'SALESFORCE_CLIENT_SECRET', label: 'Client secret', type: 'text', required: true }, + { key: 'SALESFORCE_REFRESH_TOKEN', label: 'Refresh token', type: 'text', required: true }, + ], + http_generic: [], + }; + + function setConfigAdminError(message) { + const el = document.getElementById('cfg-error'); + if (!el) return; + if (message) { + el.textContent = message; + el.classList.remove('hidden'); + } else { + el.textContent = ''; + el.classList.add('hidden'); + } + } + + function renderSecretFields(connectorId, existingKeys = []) { + const host = document.getElementById('cfg-secret-fields'); + if (!host) return; + host.innerHTML = ''; + const fields = _connectorSecretFields[connectorId] || []; + const existing = new Set(existingKeys || []); + if (!fields.length) { + host.innerHTML = '

No credential fields for this connector.

'; + return; + } + fields.forEach((f) => { + const wrap = document.createElement('div'); + wrap.className = 'field-group'; + const alreadySet = existing.has(f.key); + const label = document.createElement('label'); + label.setAttribute('for', `cfg-secret-${f.key}`); + label.textContent = + f.label + (f.required && !alreadySet ? ' *' : '') + (alreadySet ? ' (saved)' : ''); + let input; + if (f.type === 'textarea') { + input = document.createElement('textarea'); + input.rows = 5; + } else { + input = document.createElement('input'); + input.type = 'text'; + } + input.id = `cfg-secret-${f.key}`; + input.dataset.secretKey = f.key; + input.dataset.alreadySet = alreadySet ? '1' : '0'; + input.autocomplete = 'off'; + input.value = ''; + if (alreadySet) { + input.placeholder = '(already set — leave blank to keep)'; + } else if (f.required) { + input.placeholder = 'Required'; + } + wrap.appendChild(label); + wrap.appendChild(input); + host.appendChild(wrap); + }); + if ([...existing].some((k) => fields.some((f) => f.key === k))) { + const hint = document.createElement('p'); + hint.className = 'field-help'; + hint.textContent = + 'Blank credential fields keep the previously saved value. Only type a field to change it.'; + host.appendChild(hint); + } + } + + function collectSecretFields() { + const host = document.getElementById('cfg-secret-fields'); + const secrets = {}; + if (!host) return secrets; + host.querySelectorAll('[data-secret-key]').forEach((el) => { + const key = el.dataset.secretKey; + const val = (el.value || '').trim(); + if (key && val) secrets[key] = val; + }); + return secrets; + } + + async function refreshConfigDropdown(modeOrConnectorId) { + if (!_multitenancyEnabled) return; + const tenantId = getPlaygroundTenantId(); + const select = document.getElementById('playground-config-name'); + if (!select) return; + const previous = select.value; + select.innerHTML = ''; + if (!tenantId) return; + + const connectorId = + _modeToConnectorId[modeOrConnectorId] || + (Object.values(_modeToConnectorId).includes(modeOrConnectorId) + ? modeOrConnectorId + : currentConnectorId()); + if (!connectorId) return; + + const headers = { 'Content-Type': 'application/json', 'X-Tenant-ID': tenantId }; + try { + const res = await fetch( + `/v1/connectors/${encodeURIComponent(connectorId)}/configs`, + { headers } + ); + if (!res.ok) return; + const data = await res.json(); + const configs = Array.isArray(data.configs) + ? data.configs + : Array.isArray(data) + ? data + : []; + let matched = false; + configs.forEach((cfg) => { + const opt = document.createElement('option'); + opt.value = cfg.name || ''; + opt.textContent = cfg.name + (cfg.default ? ' (default)' : ''); + select.appendChild(opt); + if (previous && opt.value === previous) matched = true; + }); + if (previous && matched) { + select.value = previous; + } else if (previous && !matched) { + // Clear stale name from another connector (e.g. Drive → Epic). + select.value = ''; + } else { + const def = configs.find((c) => c.default); + if (def) select.value = def.name || ''; + } + } catch (_) { + // ignore + } + } + + function getPlaygroundTenantId() { + if (!_multitenancyEnabled) return ''; + const el = document.getElementById('playground-tenant-id'); + return (el && el.value ? el.value : '').trim(); + } + + function getPlaygroundConfigName() { + if (!_multitenancyEnabled) return ''; + const el = document.getElementById('playground-config-name'); + return (el && el.value ? el.value : '').trim(); + } + + loadFeatureFlags(); + + function logTenancyContext(actionLabel) { + if (!_multitenancyEnabled) return; + const tenantId = getPlaygroundTenantId(); + const configName = getPlaygroundConfigName(); + const label = actionLabel ? `${actionLabel} — ` : ''; + log( + `${label}Tenant: ${tenantId || '(missing)'} | Config: ${configName || '(default)'}`, + 'system' + ); + } + + function playgroundRequestHeaders(extra = {}) { + const headers = { 'Content-Type': 'application/json', ...extra }; + if (!_multitenancyEnabled) return headers; + const tenantId = getPlaygroundTenantId(); + if (tenantId) { + headers['X-Tenant-ID'] = tenantId; + } + return headers; + } + + function withTenancyQuery(endpoint) { + if (!_multitenancyEnabled) return endpoint; + const configName = getPlaygroundConfigName(); + if (!configName) return endpoint; + const sep = endpoint.includes('?') ? '&' : '?'; + return `${endpoint}${sep}config_name=${encodeURIComponent(configName)}`; + } + async function handleSubmission(payload, endpoint, btn, btnLbl, spinner, resetText, pipelineLabelOverride = null) { resetUI(pipelineLabelOverride); @@ -787,22 +1374,36 @@ document.addEventListener('DOMContentLoaded', () => { btnLbl.textContent = 'Orchestrating...'; log(`Initiating intelligent ${currentMode.toUpperCase()} orchestration...`, 'system'); + logTenancyContext(currentMode.toUpperCase()); + const configName = getPlaygroundConfigName(); + if (configName) { + payload = { ...payload, config_name: configName }; + } + const firstNode = nodes.find((n) => n && !n.classList.contains('hidden')); if (firstNode) firstNode.classList.add('active'); try { - const response = await fetch(endpoint, { + const response = await fetch(withTenancyQuery(endpoint), { method: 'POST', - headers: { 'Content-Type': 'application/json' }, + headers: playgroundRequestHeaders(), body: JSON.stringify(payload) }); - if (!response.ok) throw new Error(`Server returned ${response.status}`); + const data = await response.json().catch(() => ({})); - const data = await response.json(); - traceDisplay.textContent = data.trace_id.toUpperCase(); + if (!response.ok) { + const detail = data.detail || data.message || `Server returned ${response.status}`; + let msg = typeof detail === 'string' ? detail : JSON.stringify(detail); + if (response.status === 403 && /No connector configuration/i.test(msg)) { + msg = `${msg} Open Add config on this connector page.`; + } + throw new Error(msg); + } - for (let i = 0; i < data.steps.length; i++) { + traceDisplay.textContent = (data.trace_id || '').toUpperCase() || '—'; + + for (let i = 0; i < (data.steps || []).length; i++) { const step = data.steps[i]; const node = nodes[i]; if (!node) continue; @@ -1250,6 +1851,135 @@ document.addEventListener('DOMContentLoaded', () => { } }); + // --- Per-connector runtime config CRUD --- + const addConfigBtn = document.getElementById('add-config-btn'); + const connectorAddConfigBtn = document.getElementById('connector-add-config-btn'); + + function modalTenantId() { + const fromModal = (_cfgTenantInput && _cfgTenantInput.value ? _cfgTenantInput.value : '').trim(); + return fromModal || getPlaygroundTenantId(); + } + + async function cfgFetchFor(connectorId, method, path = '', body = null, tenantOverride = null) { + const tenantId = (tenantOverride || getPlaygroundTenantId() || '').trim(); + if (!tenantId) { + throw new Error('Set Tenant ID before managing configs (e.g. acme).'); + } + const opts = { + method, + headers: { 'Content-Type': 'application/json', 'X-Tenant-ID': tenantId }, + }; + if (body != null) { + opts.body = JSON.stringify(body); + } + const response = await fetch( + `/v1/connectors/${encodeURIComponent(connectorId)}/configs${path}`, + opts + ); + const data = await response.json().catch(() => ({})); + if (!response.ok) { + const detail = data.detail || `HTTP ${response.status}`; + throw new Error(typeof detail === 'string' ? detail : JSON.stringify(detail)); + } + return data; + } + + if (addConfigBtn) { + addConfigBtn.addEventListener('click', () => { + openConfigAdminModal(); + }); + } + if (connectorAddConfigBtn) { + connectorAddConfigBtn.addEventListener('click', () => { + openConfigAdminModal(); + }); + } + + document.getElementById('config-admin-close')?.addEventListener('click', () => { + closeConfigAdminModal(); + }); + + document.getElementById('config-admin-modal')?.addEventListener('click', (ev) => { + if (ev.target && ev.target.hasAttribute('data-config-admin-dismiss')) { + closeConfigAdminModal(); + } + }); + + document.addEventListener('keydown', (ev) => { + if (ev.key !== 'Escape') return; + const modal = document.getElementById('config-admin-modal'); + if (modal && !modal.classList.contains('hidden')) { + closeConfigAdminModal(); + } + }); + + document.getElementById('cfg-create')?.addEventListener('click', async () => { + setConfigAdminError(''); + try { + const connectorId = currentConnectorId(); + if (!connectorId) throw new Error('Open a connector page first.'); + const tenantId = modalTenantId(); + if (!tenantId) throw new Error('Tenant ID is required'); + const name = modalSelectedConfigName(); + if (!name) throw new Error('Config name is required'); + const isDefault = document.getElementById('cfg-default')?.value === 'true'; + const secrets = collectSecretFields(); + const tmpl = _connectorCreateDocs[connectorId] || { config: {}, auth: {} }; + const body = { + name, + default: isDefault, + config: tmpl.config || {}, + auth: tmpl.auth || {}, + secrets, + }; + logTenancyContext(`Config SAVE ${connectorId}`); + await cfgFetchFor(connectorId, 'POST', '', body, tenantId); + ensureTenantOption(tenantId, true); + await refreshTenantDropdown(); + ensureTenantOption(tenantId, true); + await refreshConfigDropdown(currentMode); + const headerCfg = document.getElementById('playground-config-name'); + if (headerCfg) headerCfg.value = name; + await refreshModalConfigSelect(tenantId, connectorId, name); + await applyModalConfigSelection(connectorId); + log(`Saved config '${name}' for ${connectorId} / tenant '${tenantId}'`, 'success'); + } catch (err) { + setConfigAdminError(err.message); + log(`Config save failed: ${err.message}`, 'error'); + } + }); + + document.getElementById('cfg-delete')?.addEventListener('click', async () => { + setConfigAdminError(''); + try { + const connectorId = currentConnectorId(); + if (!connectorId) throw new Error('Open a connector page first.'); + const tenantId = modalTenantId(); + const name = modalSelectedConfigName(); + if (!name) throw new Error('Config name is required'); + const newDefault = name === 'default' ? '' : 'default'; + const q = newDefault ? `?new_default=${encodeURIComponent(newDefault)}` : ''; + await cfgFetchFor( + connectorId, + 'DELETE', + `/${encodeURIComponent(name)}${q}`, + null, + tenantId + ); + log(`Deleted config '${name}' for ${connectorId} / tenant '${tenantId}'`, 'success'); + const configSelect = document.getElementById('playground-config-name'); + if (configSelect && configSelect.value === name) { + configSelect.value = ''; + } + await refreshConfigDropdown(currentMode); + await refreshModalConfigSelect(tenantId, connectorId, ''); + await applyModalConfigSelection(connectorId); + } catch (err) { + setConfigAdminError(err.message); + log(`Config delete failed: ${err.message}`, 'error'); + } + }); + if (slackActionSelect) { slackActionSelect.addEventListener('change', () => { const action = slackActionSelect.value; @@ -1557,9 +2287,13 @@ document.addEventListener('DOMContentLoaded', () => { } startRunningTimer(); - const response = await fetch('/scenarios/agent-chat-stream', { + if (_multitenancyEnabled && !getPlaygroundTenantId()) { + throw new Error('Tenant ID is required when multitenancy is enabled'); + } + logTenancyContext('Agent chat (stream)'); + const response = await fetch(withTenancyQuery('/scenarios/agent-chat-stream'), { method: 'POST', - headers: { 'Content-Type': 'application/json' }, + headers: playgroundRequestHeaders(), body: JSON.stringify({ message: message, history: agentConversationHistory.slice(0, -1) @@ -1637,9 +2371,13 @@ document.addEventListener('DOMContentLoaded', () => { return; } - const response = await fetch('/scenarios/agent-chat', { + if (_multitenancyEnabled && !getPlaygroundTenantId()) { + throw new Error('Tenant ID is required when multitenancy is enabled'); + } + logTenancyContext('Agent chat'); + const response = await fetch(withTenancyQuery('/scenarios/agent-chat'), { method: 'POST', - headers: { 'Content-Type': 'application/json' }, + headers: playgroundRequestHeaders(), body: JSON.stringify({ message: message, history: agentConversationHistory.slice(0, -1) // Exclude current message (already in payload) diff --git a/playground/index.html b/playground/index.html index 0882d126..6bf17cf4 100644 --- a/playground/index.html +++ b/playground/index.html @@ -54,6 +54,23 @@

node-wire

Auto-retry Enabled + + @@ -129,10 +146,12 @@

External Patient Viewer

- +
+ +
@@ -292,10 +311,13 @@

Salesforce

+ + + + diff --git a/playground/scenarios.py b/playground/scenarios.py index dd6d0b25..50b0f97e 100644 --- a/playground/scenarios.py +++ b/playground/scenarios.py @@ -11,7 +11,7 @@ from datetime import datetime, timezone from typing import Any, Dict, List, Optional -from fastapi import APIRouter, Depends, HTTPException +from fastapi import APIRouter, Depends, HTTPException, Request from fastapi.responses import StreamingResponse from pydantic import BaseModel, ValidationError, model_validator from dotenv import load_dotenv @@ -19,6 +19,13 @@ import asyncio from node_wire_runtime.errors import ErrorMapper from node_wire_runtime.models import ErrorCategory +from node_wire_runtime.config_store import ConfigNotFoundError +from node_wire_runtime.identity import ( + MissingTenantError, + is_multitenancy_enabled, + resolve_config_name, + resolve_tenant_id, +) from node_wire_fhir_epic.logic import FhirEpicConnector from node_wire_fhir_epic.schema import ( FhirDocumentReferenceCreateInput, @@ -73,6 +80,8 @@ logger = logging.getLogger("playground.scenarios") + + router = APIRouter(prefix="/scenarios", tags=["scenarios"]) @@ -153,6 +162,8 @@ class GoogleDriveArchivalInput(BaseModel): update_mime_type: Optional[str] = None update_add_parents: Optional[str] = None update_remove_parents: Optional[str] = None + # Resolution-time only (not a Drive API field); optional named config for the tenant. + config_name: Optional[str] = None @model_validator(mode="after") def require_upload_fields_when_not_list(self) -> "GoogleDriveArchivalInput": @@ -314,76 +325,111 @@ async def execute_with_retry( raise last_exception -# Single shared factory for playground scenarios (matches REST: enabled + exposed_via includes "rest"). +# Share the REST binding factory so playground config CRUD (/v1/...) and scenarios +# see the same in-memory ConnectorConfigStore. _playground_factory: Optional[Any] = None def get_playground_factory() -> Any: - """Lazily load connector config once; same pattern as bindings REST `get_factory`.""" + """Reuse the REST API factory (same process, same config store).""" global _playground_factory if _playground_factory is None: - from bindings.factory import ConnectorFactory - from node_wire_runtime.connector_registry import auto_register + from bindings.rest_api.app import get_factory - _playground_factory = ConnectorFactory() - auto_register() - _playground_factory.load() + _playground_factory = get_factory() return _playground_factory -def resolve_connector(connector_id: str, action: Optional[str] = None) -> Any: - """Resolve a connector via public factory API (protocol-aware).""" - factory = get_playground_factory() - return factory.get_for_protocol(connector_id, "rest", action=action) +async def resolve_connector( + request: Request, + connector_id: str, + *, + action: Optional[str] = None, + config_name: Optional[str] = None, +) -> Any: + """Resolve a tenant-scoped connector instance for the playground. + Tenant comes from ``X-Tenant-ID`` (then JWT). When multitenancy is enabled, + a tenant id is required. ``config_name`` may be supplied as a query param + or (for GDrive) payload field. + """ + from bindings.rest_api.auth import get_rest_caller_identity -def get_fhir_connector() -> FhirEpicConnector: - connector = resolve_connector("fhir_epic") - if not connector: - raise HTTPException(status_code=500, detail="FHIR Epic connector not configured") + factory = get_playground_factory() + if not factory.is_exposed(connector_id, "rest"): + raise HTTPException( + status_code=500, detail=f"Connector {connector_id!r} is not configured for REST" + ) + + try: + tenant_id = resolve_tenant_id( + headers=request.headers, + jwt_identity=get_rest_caller_identity(request), + ) + except MissingTenantError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + name = resolve_config_name((config_name or "").strip() or None) + if is_multitenancy_enabled(): + # Inline replace so CodeQL treats newline stripping as a sanitizer. + logger.info( + "Playground action | connector=%s | tenant_id=%s | config_name=%s", + str(connector_id).replace("\r", " ").replace("\n", " "), + str(tenant_id).replace("\r", " ").replace("\n", " "), + str(name or "(default)").replace("\r", " ").replace("\n", " "), + ) + try: + return await factory.get( + connector_id, + tenant_id=tenant_id, + config_name=name, + action=action, + ) + except ConfigNotFoundError: + raise HTTPException( + status_code=403, + detail="No connector configuration for this tenant", + ) + except ValueError as exc: + # Incomplete auth blocks (e.g. service_account without sa_json_secret). + raise HTTPException(status_code=400, detail=str(exc)) from exc + + +async def get_fhir_connector( + request: Request, config_name: Optional[str] = None +) -> FhirEpicConnector: + connector = await resolve_connector(request, "fhir_epic", config_name=config_name) return connector # type: ignore[return-value] -def get_http_connector(): - # Manifest action for http_generic is "request"; pass it for parity with REST routing. - connector = resolve_connector("http_generic", action="request") - if not connector: - raise HTTPException(status_code=500, detail="Generic HTTP connector not configured") +async def get_http_connector(request: Request, config_name: Optional[str] = None): + connector = await resolve_connector( + request, "http_generic", action="request", config_name=config_name + ) return connector -def get_cerner_connector(): - connector = resolve_connector("fhir_cerner") - if not connector: - raise HTTPException(status_code=500, detail="FHIR Cerner connector not configured") +async def get_cerner_connector(request: Request, config_name: Optional[str] = None): + connector = await resolve_connector(request, "fhir_cerner", config_name=config_name) return connector -def get_google_drive_connector(): - connector = resolve_connector("google_drive") - if not connector: - raise HTTPException(status_code=500, detail="Google Drive connector not configured") - return connector +async def get_google_drive_connector(request: Request, config_name: Optional[str] = None) -> Any: + """``config_name`` query param is optional; GDrive body may also carry it.""" + return await resolve_connector(request, "google_drive", config_name=config_name) -def get_slack_connector(): - connector = resolve_connector("slack") - if not connector: - raise HTTPException(status_code=500, detail="Slack connector not configured") +async def get_slack_connector(request: Request, config_name: Optional[str] = None): + connector = await resolve_connector(request, "slack", config_name=config_name) return connector -def get_stripe_connector(): - connector = resolve_connector("stripe") - if not connector: - raise HTTPException(status_code=500, detail="Stripe connector not configured") +async def get_stripe_connector(request: Request, config_name: Optional[str] = None): + connector = await resolve_connector(request, "stripe", config_name=config_name) return connector -def get_salesforce_connector(): - connector = resolve_connector("salesforce") - if not connector: - raise HTTPException(status_code=500, detail="Salesforce connector not configured") +async def get_salesforce_connector(request: Request, config_name: Optional[str] = None): + connector = await resolve_connector(request, "salesforce", config_name=config_name) return connector @@ -1348,9 +1394,15 @@ def add_step( @router.post("/gdrive-archival", response_model=ScenarioResponse) async def gdrive_archival_scenario( - payload: GoogleDriveArchivalInput, connector: Any = Depends(get_google_drive_connector) + request: Request, + payload: GoogleDriveArchivalInput, ) -> ScenarioResponse: - """4-step Google Drive archival and sharing demo.""" + """4-step Google Drive archival and sharing demo (tenant-aware).""" + connector = await resolve_connector( + request, + "google_drive", + config_name=payload.config_name, + ) trace_id = str(uuid.uuid4()) steps: List[ScenarioStep] = [] @@ -1845,15 +1897,47 @@ def _current_agent_transport() -> str: return transport if transport in {"stdio", "streamable-http"} else "stdio" +def _resolve_agent_mcp_tenant(request: Request) -> str: + """Resolve tenant for playground agent → MCP (stdio pin / HTTP header).""" + from bindings.rest_api.auth import get_rest_caller_identity + + try: + return resolve_tenant_id( + headers=request.headers, + jwt_identity=get_rest_caller_identity(request), + ) + except MissingTenantError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + + +def _agent_mcp_extra_headers(tenant_id: str) -> Dict[str, str]: + from node_wire_runtime.identity import TENANT_HEADER + + # TENANT_HEADER is lowercased; HTTP clients accept any casing. + return {TENANT_HEADER: tenant_id} + + +def _playground_inprocess_mcp_client(tenant_id: str, request: Request): + """Share playground factory/store with agent tool calls (named tenants).""" + from agents.toolhive import InProcessMcpClient + from bindings.mcp_server.server import McpServer + + config_name = resolve_config_name( + (request.query_params.get("config_name") or "").strip() or None + ) + factory = get_playground_factory() + server = McpServer(server_name="node-wire-playground-agent", factory=factory) + return InProcessMcpClient(server, tenant_id=tenant_id, config_name=config_name) + + @router.post("/agent-chat", response_model=AgentChatResponse) -async def agent_chat(payload: AgentChatInput) -> AgentChatResponse: +async def agent_chat(request: Request, payload: AgentChatInput) -> AgentChatResponse: """ AI Agent chatbot endpoint. Accepts a user message + conversation history, runs through the ToolHiveAgent, and returns the agent's reply with any tool steps executed. """ import os - import sys trace_id = str(uuid.uuid4()) logger.info( @@ -1870,12 +1954,23 @@ async def agent_chat(payload: AgentChatInput) -> AgentChatResponse: success=False, ) + tenant_id = _resolve_agent_mcp_tenant(request) + mcp_headers = _agent_mcp_extra_headers(tenant_id) + config_name = resolve_config_name( + (request.query_params.get("config_name") or "").strip() or None + ) + if is_multitenancy_enabled(): + logger.info( + "Agent Chat | mcp_tenant_id=%s | config_name=%s", + str(tenant_id).replace("\r", " ").replace("\n", " "), + str(config_name or "(default)").replace("\r", " ").replace("\n", " "), + ) + try: from agents.llm_factory import LLMProviderFactory from agents.toolhive import ( MultiMcpClient, ToolHiveAgent, - StdioMcpClient, resolve_mcp_urls, resolve_max_tool_failures, ) @@ -1887,21 +1982,23 @@ async def agent_chat(payload: AgentChatInput) -> AgentChatResponse: task = _build_agent_chat_task(payload) - # Determine MCP transport — try proxy first, fallback to local stdio + # Determine MCP transport — try proxy first, fallback to local in-process transport = _current_agent_transport() urls = resolve_mcp_urls() if transport == "streamable-http" else [] run_result = None fallback_to_stdio = ( - os.environ.get("PLAYGROUND_AGENT_PROXY_FALLBACK_TO_STDIO", "false").lower() == "true" + os.environ.get("PLAYGROUND_AGENT_PROXY_FALLBACK_TO_STDIO", "true").lower() == "true" ) if urls: logger.info("Agent Chat | trying ToolHive proxy URL(s): %s", ",".join(urls)) try: if len(urls) == 1: - mcp_client = create_http_mcp_client(urls[0]) + mcp_client = create_http_mcp_client(urls[0], extra_headers=mcp_headers) else: - mcp_client = MultiMcpClient([create_http_mcp_client(u) for u in urls]) + mcp_client = MultiMcpClient( + [create_http_mcp_client(u, extra_headers=mcp_headers) for u in urls] + ) agent = ToolHiveAgent( mcp_client, llm_provider, @@ -1951,10 +2048,9 @@ async def agent_chat(payload: AgentChatInput) -> AgentChatResponse: ) if run_result is None: - # Use local stdio transport - logger.info("Agent Chat | using local stdio MCP transport") - cmd = [sys.executable, "-m", "agents.mcp_entrypoint"] - async with StdioMcpClient(cmd) as mcp_client: + # In-process MCP shares the playground config store (named tenants). + logger.info("Agent Chat | using in-process MCP (shared factory)") + async with _playground_inprocess_mcp_client(tenant_id, request) as mcp_client: agent = ToolHiveAgent( mcp_client, llm_provider, @@ -1999,22 +2095,30 @@ async def agent_chat(payload: AgentChatInput) -> AgentChatResponse: @router.post("/agent-chat-stream") -async def agent_chat_stream(payload: AgentChatInput) -> Any: +async def agent_chat_stream(request: Request, payload: AgentChatInput) -> Any: """ Stream agent progress and final-answer chunks to web clients. The terminal ``done`` event includes ``trace_id`` and ``message``. Clients should stop their streaming loader only when that event arrives. """ + tenant_id = _resolve_agent_mcp_tenant(request) + mcp_headers = _agent_mcp_extra_headers(tenant_id) + config_name = resolve_config_name( + (request.query_params.get("config_name") or "").strip() or None + ) + if is_multitenancy_enabled(): + logger.info( + "Agent Chat stream | mcp_tenant_id=%s | config_name=%s", + str(tenant_id).replace("\r", " ").replace("\n", " "), + str(config_name or "(default)").replace("\r", " ").replace("\n", " "), + ) async def stream_events(): try: - import sys - from agents.llm_factory import LLMProviderFactory from agents.toolhive import ( MultiMcpClient, - StdioMcpClient, ToolHiveAgent, resolve_mcp_urls, resolve_max_tool_failures, @@ -2049,25 +2153,55 @@ async def stream_events(): task = _build_agent_chat_task(payload) transport = _current_agent_transport() urls = resolve_mcp_urls() if transport == "streamable-http" else [] + # Default true: playground should work without a separate MCP on :8081. + fallback_to_local = ( + os.environ.get("PLAYGROUND_AGENT_PROXY_FALLBACK_TO_STDIO", "true").lower() == "true" + ) if urls: - if len(urls) == 1: - mcp_client = create_http_mcp_client(urls[0]) - else: - mcp_client = MultiMcpClient([create_http_mcp_client(u) for u in urls]) - agent = ToolHiveAgent( - mcp_client, - llm_provider, - max_steps=10, - max_tool_failures=resolve_max_tool_failures(None), - ) - agent._system_prompt = AGENT_GUARDRAIL_PROMPT - async for event in agent.run_events(task): - yield json.dumps(event) + "\n" - return + proxy_ok = False + try: + if len(urls) == 1: + mcp_client = create_http_mcp_client(urls[0], extra_headers=mcp_headers) + else: + mcp_client = MultiMcpClient( + [create_http_mcp_client(u, extra_headers=mcp_headers) for u in urls] + ) + agent = ToolHiveAgent( + mcp_client, + llm_provider, + max_steps=10, + max_tool_failures=resolve_max_tool_failures(None), + ) + agent._system_prompt = AGENT_GUARDRAIL_PROMPT + async for event in agent.run_events(task): + if ( + fallback_to_local + and event.get("type") == "error" + and "Failed to list MCP tools" in str(event.get("message") or "") + ): + logger.warning( + "Agent Chat stream | proxy incomplete (%s) — " + "falling back to in-process MCP", + event.get("message"), + ) + break + proxy_ok = True + yield json.dumps(event) + "\n" + else: + # Exhausted generator without break → proxy finished normally. + return + if proxy_ok: + return + except Exception as proxy_err: + if not fallback_to_local: + raise + logger.warning( + "Agent Chat stream | proxy error: %s — falling back to in-process MCP", + proxy_err, + ) - cmd = [sys.executable, "-m", "agents.mcp_entrypoint"] - async with StdioMcpClient(cmd) as mcp_client: + async with _playground_inprocess_mcp_client(tenant_id, request) as mcp_client: agent = ToolHiveAgent( mcp_client, llm_provider, diff --git a/playground/style.css b/playground/style.css index 934f763a..9298055b 100644 --- a/playground/style.css +++ b/playground/style.css @@ -99,6 +99,200 @@ body { align-items: center; } +.header-tenancy { + display: flex; + gap: 0.75rem; + align-items: flex-end; +} + +.header-tenancy-field { + display: flex; + flex-direction: column; + gap: 0.2rem; +} + +.header-tenancy-field label { + font-size: 0.65rem; + font-weight: 600; + text-transform: uppercase; + letter-spacing: 0.04em; + color: var(--text-muted); +} + +.header-tenancy-field input, +.header-tenancy-field select { + background: rgba(255, 255, 255, 0.06); + border: 1px solid var(--border); + border-radius: 8px; + color: var(--text-main); + font-size: 0.8rem; + padding: 0.35rem 0.55rem; + min-width: 7.5rem; + max-width: 10rem; +} + +.header-tenancy-field input:focus, +.header-tenancy-field select:focus { + outline: none; + border-color: var(--brand-accent); +} + +.connectors-toolbar { + display: flex; + gap: 0.75rem; + align-items: center; + justify-content: flex-start; + margin-bottom: 1.5rem; + flex-wrap: wrap; +} + +.connectors-toolbar .main-back-btn { + margin-bottom: 0; +} + +.add-config-btn { + display: inline-flex; + align-items: center; + justify-content: center; + gap: 0.35rem; + align-self: flex-end; + min-height: 2.05rem; + padding: 0.35rem 0.85rem; + border-radius: 8px; + border: 1px solid color-mix(in srgb, var(--brand-accent) 55%, var(--border)); + background: color-mix(in srgb, var(--brand-accent) 18%, transparent); + color: var(--text-main); + font-family: inherit; + font-size: 0.8rem; + font-weight: 650; + letter-spacing: 0.01em; + cursor: pointer; + white-space: nowrap; + box-shadow: 0 1px 0 rgba(255, 255, 255, 0.06) inset; + transition: background 0.2s, border-color 0.2s, transform 0.2s, box-shadow 0.2s; +} + +.add-config-btn:hover { + background: color-mix(in srgb, var(--brand-accent) 28%, transparent); + border-color: var(--brand-accent); + transform: translateY(-1px); + box-shadow: 0 6px 18px rgba(0, 0, 0, 0.18); +} + +.add-config-btn:focus-visible { + outline: 2px solid var(--brand-accent); + outline-offset: 2px; +} + +.config-admin-modal { + position: fixed; + inset: 0; + z-index: 1200; + display: flex; + align-items: center; + justify-content: center; + padding: 1.25rem; +} + +.config-admin-modal.hidden { + display: none; +} + +.config-admin-modal-backdrop { + position: absolute; + inset: 0; + background: rgba(15, 23, 42, 0.55); + backdrop-filter: blur(2px); +} + +.config-admin-modal-card { + position: relative; + z-index: 1; + width: min(40rem, 100%); + max-height: min(90vh, 52rem); + overflow: auto; + margin: 0; + padding: 1.35rem 1.6rem 1.5rem; + box-shadow: 0 24px 48px rgba(0, 0, 0, 0.35); +} + +.config-admin-modal-title { + display: flex; + align-items: flex-start; + justify-content: space-between; + gap: 1rem; + margin-bottom: 0.65rem; +} + +.config-admin-close { + display: inline-flex; + align-items: center; + justify-content: center; + flex-shrink: 0; + width: 2rem; + height: 2rem; + border: 1px solid var(--border); + border-radius: 8px; + background: transparent; + color: var(--text-muted); + cursor: pointer; + transition: background 0.15s, color 0.15s, border-color 0.15s; +} + +.config-admin-close:hover { + color: var(--text-main); + border-color: var(--brand-accent); + background: color-mix(in srgb, var(--brand-accent) 12%, transparent); +} + +.config-admin-panel { + padding: 1.35rem 1.6rem 1.5rem; +} + +.config-admin-panel .card-title { + margin-bottom: 0.65rem; +} + +.config-admin-help { + margin: 0 0 1rem; + max-width: 52rem; + line-height: 1.45; +} + +.config-admin-fields { + gap: 1rem; + margin-bottom: 0.25rem; +} + +.config-admin-secret-fields { + display: flex; + flex-direction: column; + gap: 0.75rem; + margin: 0.75rem 0 1rem; +} + +.config-admin-secret-fields .field-group textarea { + min-height: 6rem; + font-family: ui-monospace, SFMono-Regular, Menlo, Consolas, monospace; + font-size: 0.8rem; +} + +.config-admin-error { + margin: 0.5rem 0 0; + color: #f87171; + font-size: 0.85rem; + line-height: 1.4; +} + +.config-admin-error.hidden { + display: none; +} + +.config-admin-fields .field-group { + flex: 1 1 12rem; + min-width: 10rem; +} + .business-value-stat { text-align: right; border-right: 1px solid #cbd5e1; @@ -267,9 +461,35 @@ body { background: rgba(37, 99, 235, 0.1); color: var(--brand-accent); padding: 0.25rem 0.75rem; - border-radius: 6px; + border-radius: 999px; font-size: 0.75rem; - font-weight: 700; + font-weight: 600; +} + +.tenancy-config-actions { + display: flex; + flex-wrap: wrap; + gap: 0.65rem; + margin-top: 1rem; +} + +.tenancy-cfg-btn { + width: auto; + min-height: 2.4rem; + padding: 0.6rem 1.05rem; + font-size: 0.875rem; + font-weight: 600; + border-radius: 10px; +} + +.tenancy-cfg-btn-danger { + background: transparent; + border: 1px solid rgba(239, 68, 68, 0.45); + color: #f87171; +} + +.tenancy-cfg-btn-danger:hover { + background: rgba(239, 68, 68, 0.12); } /* Form Styles */ @@ -1098,6 +1318,27 @@ input[type="range"]::-webkit-slider-thumb:hover { line-height: 1.6; } +/* Connector page top bar: back + Add config */ +.connector-page-toolbar { + display: flex; + align-items: center; + justify-content: space-between; + gap: 0.75rem 1rem; + flex-wrap: wrap; + margin: 0 0 1.25rem; + min-height: 2.25rem; +} + +.connector-page-toolbar .back-link { + margin-bottom: 0; +} + +.connector-page-toolbar .add-config-btn { + align-self: center; + margin: 0; + flex-shrink: 0; +} + /* Back Link Navigation */ .back-link { display: flex; diff --git a/pyproject.toml b/pyproject.toml index 5c297603..bbe8674e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -125,7 +125,7 @@ override-dependencies = [ ] [tool.pytest.ini_options] -pythonpath = ["src", "."] +pythonpath = ["src", ".", "src/bindings/grpc_server"] asyncio_mode = "auto" asyncio_default_fixture_loop_scope = "function" addopts = [ diff --git a/sample.env b/sample.env index 389e5c98..b880d8ca 100644 --- a/sample.env +++ b/sample.env @@ -2,6 +2,60 @@ # # SPDX-License-Identifier: Apache-2.0 +# Multi-tenancy (default: false — single-tenant / legacy mode) +# Set to true to enable header-based tenant routing (X-Tenant-ID), per-tenant +# named configs, and tenancy controls in the playground. +NW_MULTITENANCY_ENABLED=false +# +# MCP stdio / ToolHive: pin one process to one tenant (ignored on streamable-http; +# HTTP clients must send X-Tenant-ID instead). Required when multitenancy is enabled +# for stdio launches. +# NW_TENANT_ID=acme +# Optional header name override (REST / MCP HTTP / gRPC metadata): +# NW_TENANT_ID_HEADER=X-Tenant-ID +# +# Named tenants (anything other than __default__) resolve secrets strictly via +# TenantSecretProvider as: NW_{TENANT}_{CONNECTOR}_{KEY} +# (segments uppercased; non-alphanumeric → _). No fallback to shared env vars. +# Example tenant id "acme" — copy/uncomment and set values (or ${SHARED_VAR} refs): +# +# # fhir_epic — auth keys from connectors.yaml + epic_fhir_base_url from connector logic +# NW_ACME_FHIR_EPIC_EPIC_FHIR_BASE_URL=${EPIC_FHIR_BASE_URL} +# NW_ACME_FHIR_EPIC_EPIC_TOKEN_URL=${EPIC_TOKEN_URL} +# NW_ACME_FHIR_EPIC_EPIC_CLIENT_ID=${EPIC_CLIENT_ID} +# NW_ACME_FHIR_EPIC_EPIC_PRIVATE_KEY=${EPIC_PRIVATE_KEY} +# NW_ACME_FHIR_EPIC_EPIC_KID=${EPIC_KID} +# +# # fhir_cerner +# NW_ACME_FHIR_CERNER_CERNER_FHIR_BASE_URL=${CERNER_FHIR_BASE_URL} +# NW_ACME_FHIR_CERNER_CERNER_TOKEN_URL=${CERNER_TOKEN_URL} +# NW_ACME_FHIR_CERNER_CERNER_CLIENT_ID=${CERNER_CLIENT_ID} +# NW_ACME_FHIR_CERNER_CERNER_PRIVATE_KEY=${CERNER_PRIVATE_KEY} +# NW_ACME_FHIR_CERNER_CERNER_KID=${CERNER_KID} +# NW_ACME_FHIR_CERNER_CERNER_SCOPES=${CERNER_SCOPES} +# +# # google_drive +# NW_ACME_GOOGLE_DRIVE_GOOGLE_DRIVE_SA_JSON=${GOOGLE_DRIVE_SA_JSON} +# +# # smtp +# NW_ACME_SMTP_SMTP_USERNAME=${SMTP_USERNAME} +# NW_ACME_SMTP_SMTP_PASSWORD=${SMTP_PASSWORD} +# +# # stripe (secret_key in yaml is stripe_api_key) +# NW_ACME_STRIPE_STRIPE_API_KEY=${STRIPE_API_KEY} +# +# # slack +# NW_ACME_SLACK_SLACK_BOT_TOKEN=${SLACK_BOT_TOKEN} +# +# # salesforce — auth keys + salesforce_instance_url from connector logic +# NW_ACME_SALESFORCE_SALESFORCE_INSTANCE_URL=${SALESFORCE_INSTANCE_URL} +# NW_ACME_SALESFORCE_SALESFORCE_TOKEN_URL=${SALESFORCE_TOKEN_URL} +# NW_ACME_SALESFORCE_SALESFORCE_CLIENT_ID=${SALESFORCE_CLIENT_ID} +# NW_ACME_SALESFORCE_SALESFORCE_CLIENT_SECRET=${SALESFORCE_CLIENT_SECRET} +# NW_ACME_SALESFORCE_SALESFORCE_REFRESH_TOKEN=${SALESFORCE_REFRESH_TOKEN} +# +# http_generic uses NoAuth — no tenant secrets required. + # Epic FHIR EPIC_FHIR_BASE_URL=https://fhir.epic.com/interconnect-fhir-oauth/api/FHIR/R4 EPIC_TOKEN_URL=https://fhir.epic.com/interconnect-fhir-oauth/oauth2/token diff --git a/src/agents/providers/groq_provider.py b/src/agents/providers/groq_provider.py index 8105efb4..3d084407 100644 --- a/src/agents/providers/groq_provider.py +++ b/src/agents/providers/groq_provider.py @@ -19,6 +19,8 @@ from typing import Any, Dict, List, cast from agents.llm_factory import BaseLLMProvider, LLMMessage, LLMResponse, ToolCall +from agents.schema_utils import openai_compatible_tool_parameters + logger = logging.getLogger("agents.providers.groq") @@ -30,7 +32,7 @@ def _mcp_tool_to_groq(tool: Dict[str, Any]) -> Dict[str, Any]: "function": { "name": tool["name"], "description": tool.get("description", ""), - "parameters": tool.get("input_schema", {"type": "object", "properties": {}}), + "parameters": openai_compatible_tool_parameters(tool.get("input_schema")), }, } diff --git a/src/agents/providers/openai_provider.py b/src/agents/providers/openai_provider.py index 29d1b500..95797eec 100644 --- a/src/agents/providers/openai_provider.py +++ b/src/agents/providers/openai_provider.py @@ -18,6 +18,7 @@ from typing import Any, Dict, List, cast from agents.llm_factory import BaseLLMProvider, LLMMessage, LLMResponse, ToolCall +from agents.schema_utils import openai_compatible_tool_parameters logger = logging.getLogger("agents.providers.openai") @@ -28,7 +29,7 @@ def _mcp_tool_to_openai(tool: Dict[str, Any]) -> Dict[str, Any]: "function": { "name": tool["name"], "description": tool.get("description", ""), - "parameters": tool.get("input_schema", {"type": "object", "properties": {}}), + "parameters": openai_compatible_tool_parameters(tool.get("input_schema")), }, } diff --git a/src/agents/schema_utils.py b/src/agents/schema_utils.py new file mode 100644 index 00000000..9ebf77e3 --- /dev/null +++ b/src/agents/schema_utils.py @@ -0,0 +1,36 @@ +# +# SPDX-FileCopyrightText: 2026 AOT Technologies +# SPDX-License-Identifier: Apache-2.0 +# +"""Shared JSON-Schema helpers for LLM provider tool definitions.""" + +from __future__ import annotations + +import copy +from typing import Any, Dict, Optional + + +def openai_compatible_tool_parameters(input_schema: Optional[Dict[str, Any]]) -> Dict[str, Any]: + """Copy a JSON Schema so optional properties also accept ``null``. + + Providers such as Groq validate model tool calls against the schema and reject + ``"field": null`` when the type is only ``"string"``. Models often emit null + for unused optional keys instead of omitting them. + """ + schema = copy.deepcopy(input_schema) if input_schema else {"type": "object", "properties": {}} + props = schema.get("properties") + if not isinstance(props, dict): + return schema + required = set(schema.get("required") or []) + for key, prop in props.items(): + if key in required or not isinstance(prop, dict): + continue + t = prop.get("type") + if t is None: + continue + if isinstance(t, list): + if "null" not in t: + prop["type"] = [*t, "null"] + elif t != "null": + prop["type"] = [t, "null"] + return schema diff --git a/src/agents/toolhive.py b/src/agents/toolhive.py index de8cb035..b2b33cdc 100644 --- a/src/agents/toolhive.py +++ b/src/agents/toolhive.py @@ -48,7 +48,17 @@ import uuid from contextlib import AsyncExitStack from dataclasses import dataclass, field -from typing import Any, AsyncIterator, Dict, List, Optional, Protocol, Union, runtime_checkable +from typing import ( + Any, + AsyncIterator, + Dict, + List, + Mapping, + Optional, + Protocol, + Union, + runtime_checkable, +) import re from dotenv import load_dotenv @@ -90,6 +100,11 @@ def _redact_tool_args_for_log(tool_name: str, args: Dict[str, Any]) -> Dict[str, return scrubbed +def omit_null_tool_args(arguments: Optional[Dict[str, Any]]) -> Dict[str, Any]: + """Drop keys whose value is ``None`` (LLM 'omit optional' often arrives as null).""" + return {k: v for k, v in dict(arguments or {}).items() if v is not None} + + def truncate_tool_result_for_llm(text: str) -> str: """ Cap tool output size sent to the LLM so providers with strict limits (e.g. Groq @@ -241,19 +256,26 @@ class ToolHiveMcpClient: response header must be forwarded in all subsequent requests. """ - def __init__(self, base_url: str) -> None: + def __init__( + self, + base_url: str, + *, + extra_headers: Optional[Mapping[str, str]] = None, + ) -> None: self._base_url = base_url.rstrip("/") self._session_id: Optional[str] = None self._initialized: bool = False self._auth_token: Optional[str] = os.environ.get( "TOOLHIVE_MCP_BEARER_TOKEN" ) or os.environ.get("TOOLHIVE_MCP_API_KEY") + self._extra_headers: Dict[str, str] = dict(extra_headers or {}) def _build_request_headers(self) -> Dict[str, str]: headers: Dict[str, str] = { "Content-Type": "application/json", "Accept": "application/json, text/event-stream", } + headers.update(self._extra_headers) if self._session_id: headers["Mcp-Session-Id"] = self._session_id # For MCP auth-gated servers, send both forms for compatibility. @@ -425,8 +447,14 @@ class StdioMcpClient: Useful for local manual testing without ToolHive. """ - def __init__(self, command: List[str]) -> None: + def __init__( + self, + command: List[str], + *, + env: Optional[Mapping[str, str]] = None, + ) -> None: self._command = command + self._env = dict(env) if env is not None else None self._exit_stack = AsyncExitStack() self._session: Any = None @@ -437,10 +465,13 @@ async def __aenter__(self) -> StdioMcpClient: except ImportError as exc: raise ImportError("mcp SDK not installed.") from exc + child_env = os.environ.copy() + if self._env: + child_env.update({k: str(v) for k, v in self._env.items() if v is not None}) params = StdioServerParameters( command=self._command[0], args=self._command[1:], - env=os.environ.copy(), + env=child_env, ) stdio_transport = await self._exit_stack.enter_async_context(stdio_client(params)) self._read, self._write = stdio_transport @@ -467,10 +498,52 @@ async def call_tool(self, name: str, arguments: Dict[str, Any]) -> str: if not self._session: raise RuntimeError("Client not initialised. Use 'async with'") resp = await self._session.call_tool(name, arguments) - parts = [c.text for c in resp.content if hasattr(c, "text")] + # Extract text content + parts = [] + for block in resp.content: + if hasattr(block, "text"): + parts.append(block.text) + else: + parts.append(str(block)) return "\n".join(parts) +class InProcessMcpClient: + """MCP client backed by an in-process :class:`McpServer` (shared factory/store). + + Used by the playground agent so named-tenant configs created via REST are + visible to tool calls without spawning a separate stdio process. + """ + + def __init__( + self, + server: Any, + *, + tenant_id: str, + config_name: Optional[str] = None, + ) -> None: + self._server = server + self._tenant_id = tenant_id + self._config_name = config_name + + async def __aenter__(self) -> InProcessMcpClient: + self._server._stdio_env_tenant_pin = self._tenant_id + return self + + async def __aexit__(self, exc_type, exc_val, exc_tb) -> None: + return None + + async def list_tools(self) -> List[Dict[str, Any]]: + return self._server.list_tools() + + async def call_tool(self, name: str, arguments: Dict[str, Any]) -> str: + args = omit_null_tool_args(arguments) + if self._config_name and not args.get("config_name"): + args["config_name"] = self._config_name + result = await self._server.invoke_tool(name, args) + return json.dumps(result, default=str) + + # --------------------------------------------------------------------------- # The Agent # --------------------------------------------------------------------------- @@ -617,7 +690,9 @@ async def run(self, task: str) -> AgentRunResult: ) try: - tool_result_str = await self._mcp.call_tool(tc.name, tc.arguments) + tool_result_str = await self._mcp.call_tool( + tc.name, omit_null_tool_args(tc.arguments) + ) logger.info( "Tool %s returned response of length: %d chars", tc.name, @@ -786,7 +861,9 @@ async def _run_events_inner(self, task: str, trace_id: str) -> AsyncIterator[Dic logger.info("Calling tool: %s | args=%s", tc.name, scrubbed_args) try: - tool_result_str = await self._mcp.call_tool(tc.name, tc.arguments) + tool_result_str = await self._mcp.call_tool( + tc.name, omit_null_tool_args(tc.arguments) + ) logger.info("Tool %s returned: %.200s", tc.name, tool_result_str) except Exception as exc: tool_result_str = f"ERROR: {exc}" diff --git a/src/bindings/factory.py b/src/bindings/factory.py index 1395d314..ccc0bb18 100644 --- a/src/bindings/factory.py +++ b/src/bindings/factory.py @@ -4,24 +4,38 @@ # from __future__ import annotations +import asyncio import logging import os import re +import threading +from collections import OrderedDict from dataclasses import dataclass from pathlib import Path -from typing import Any, Dict, List, Optional +from typing import Any, Dict, List, Optional, Tuple import yaml from node_wire_runtime import BaseConnector, SecretProvider, get_connector_registry -from node_wire_runtime.policy import PolicyHook +from node_wire_runtime.config_store import ( + DEFAULT_TENANT, + ConfigNotFoundError, + ConfigRecord, + ConnectorConfigStore, +) +from node_wire_runtime.policy import PolicyHook, TenantConfigHook from node_wire_runtime.policies.mcp_scope_policy import ( DEFAULT_SCOPE_MODE_DENY, ScopePolicyHook, load_scope_map_from_env, load_scope_policy_default_from_env, ) -from node_wire_runtime.secrets import ChainedSecretProvider, EnvSecretProvider +from node_wire_runtime.secrets import ( + ChainedSecretProvider, + EnvSecretProvider, + OverlaySecretProvider, + TenantSecretProvider, +) logger = logging.getLogger("bindings.factory") @@ -98,9 +112,11 @@ def _build_secret_provider() -> SecretProvider: ``NW_AWS_SECRETS_MANAGER_SECRET_ID`` — Secrets Manager secret id or ARN ``AWS_REGION`` — optional, default ``us-east-1`` """ + # Simplified: tenant secret overlay sits in front of env (and aws_env chain). + overlay = OverlaySecretProvider.instance() mode = os.environ.get("NW_SECRET_BACKEND", "env").strip().lower() if mode in ("", "env"): - return EnvSecretProvider() + return ChainedSecretProvider(overlay, EnvSecretProvider()) if mode == "aws_env": secret_id = os.environ.get("NW_AWS_SECRETS_MANAGER_SECRET_ID") if not secret_id: @@ -111,6 +127,7 @@ def _build_secret_provider() -> SecretProvider: region = os.environ.get("AWS_REGION", "us-east-1") return ChainedSecretProvider( + overlay, AwsSecretsManagerProvider(secret_name=secret_id, region=region), EnvSecretProvider(), ) @@ -171,11 +188,36 @@ class ConnectorFactory: def __init__(self, config_path: str | Path | None = None) -> None: self._config_path = _resolve_config_path(config_path) self._secret_provider: SecretProvider = _build_secret_provider() - self._policy_hook: PolicyHook | None = _build_policy_hook() - self._connectors: Dict[str, Any] = {} + # Runtime config store + per-(tenant, connector, config_name) invoke cache. + self._store = ConnectorConfigStore() + self._store.attach_factory(self) + # Scope policy takes precedence when configured; otherwise fall back to the + # config-existence hook (defense in depth for configured = entitled). + self._policy_hook: PolicyHook | None = _build_policy_hook() or TenantConfigHook(self._store) + # YAML-derived metadata (enabled, exposed_via, raw) for enumeration, + # protocol gating, and the MCP upstream-passthrough check. self._configs: Dict[str, ConnectorConfig] = {} + self._instances: "OrderedDict[Tuple[str, str, str], BaseConnector]" = OrderedDict() + # Guards all OrderedDict mutations. asyncio.Lock is per-key single-flight only + # and does not protect sync/off-loop callers (invalidate, list_for_protocol). + self._instances_guard = threading.RLock() + self._locks: Dict[Tuple[str, str, str], asyncio.Lock] = {} + self._locks_guard = threading.Lock() + self._max_instances = int(os.environ.get("NW_FACTORY_MAX_INSTANCES", "512")) + self._loop: asyncio.AbstractEventLoop | None = None + + @property + def store(self) -> ConnectorConfigStore: + """The runtime connector config store backing this factory.""" + return self._store def load(self) -> None: + """Load connector metadata and bootstrap the store from ``connectors.yaml``. + + The YAML remains the single-tenant bootstrap: it is translated once into the + store under the ``__default__`` tenant (gated by ``NW_CONFIG_BOOTSTRAP_YAML``, + default on), so existing deployments keep working with zero changes. + """ logger.info("Loading connector configuration", extra={"config_path": self._config_path}) path = Path(self._config_path) if not path.is_file(): @@ -188,6 +230,14 @@ def load(self) -> None: raw = _resolve_env_vars(raw) connectors_cfg: Dict[str, Any] = raw.get("connectors", {}) + bootstrap_enabled = os.environ.get( + "NW_CONFIG_BOOTSTRAP_YAML", "true" + ).strip().lower() not in ( + "0", + "false", + "no", + ) + bootstrap_payload: Dict[str, Any] = {DEFAULT_TENANT: {}} for connector_id, cfg in connectors_cfg.items(): enabled = bool(cfg.get("enabled", False)) @@ -220,12 +270,30 @@ def load(self) -> None: ) continue - instance = self._instantiate(connector_id) - self._connectors[connector_id] = instance + if bootstrap_enabled: + doc: Dict[str, Any] = { + "name": "default", + "default": True, + "config": { + k: v + for k, v in cfg_raw.items() + if k not in ("enabled", "exposed_via", "auth") + }, + "auth": cfg_raw.get("auth", {}), + "exposed_via": exposed_via, + } + bootstrap_payload[DEFAULT_TENANT][connector_id] = [doc] + + if bootstrap_enabled: + self._store.init(bootstrap_payload) - def _build_auth_provider(self, connector_id: str, cfg: dict) -> Any: - """Construct the appropriate AuthProvider from the connector's YAML ``auth:`` block. + def _build_auth_provider( + self, connector_id: str, cfg: dict, *, secret_provider: SecretProvider | None = None + ) -> Any: + """Construct the appropriate AuthProvider from the connector's ``auth:`` block. + ``secret_provider`` defaults to the factory's shared provider (used for + enumeration instances); the invoke path passes a tenant-scoped provider. Falls back to :class:`NoAuthProvider` when the block is absent. """ from node_wire_runtime.auth import ( @@ -235,6 +303,8 @@ def _build_auth_provider(self, connector_id: str, cfg: dict) -> Any: StaticTokenAuthProvider, ) + sp = secret_provider if secret_provider is not None else self._secret_provider + auth_cfg = cfg.get("auth") or {} if connector_id == "google_drive": auth_cfg = _resolve_google_drive_auth(auth_cfg) @@ -245,7 +315,7 @@ def _build_auth_provider(self, connector_id: str, cfg: dict) -> Any: if provider_type == "static_token": return StaticTokenAuthProvider( - secret_provider=self._secret_provider, + secret_provider=sp, secret_key=auth_cfg["secret_key"], header_name=auth_cfg.get("header_name", "Authorization"), prefix=auth_cfg.get("prefix", "Bearer"), @@ -254,7 +324,7 @@ def _build_auth_provider(self, connector_id: str, cfg: dict) -> Any: if provider_type == "oauth2": return OAuth2AuthProvider( - secret_provider=self._secret_provider, + secret_provider=sp, grant_method=auth_cfg.get("grant_method", "private_key_jwt"), token_url_secret=auth_cfg["token_url_secret"], client_id_secret=auth_cfg["client_id_secret"], @@ -271,8 +341,14 @@ def _build_auth_provider(self, connector_id: str, cfg: dict) -> Any: ) if provider_type == "service_account": + if "sa_json_secret" not in auth_cfg: + raise ValueError( + f"google_drive auth provider 'service_account' requires " + f"'sa_json_secret' in the config auth block " + f"(connector={connector_id!r})" + ) return ServiceAccountAuthProvider( - secret_provider=self._secret_provider, + secret_provider=sp, sa_json_secret=auth_cfg["sa_json_secret"], scopes=auth_cfg.get("scopes"), ) @@ -306,14 +382,17 @@ async def get_client_credentials(self): # type: ignore[override] password_secret = auth_cfg.get("password_secret", "SMTP_PASSWORD") from node_wire_runtime.auth.base import AuthProvider - sp = self._secret_provider + creds_sp = sp class _SmtpCredentialsProvider(AuthProvider): # type: ignore[misc] async def get_headers(self) -> dict: return {} async def get_client_credentials(self): # type: ignore[override] - return (sp.get_secret(username_secret), sp.get_secret(password_secret)) + return ( + creds_sp.get_secret(username_secret), + creds_sp.get_secret(password_secret), + ) return _SmtpCredentialsProvider() @@ -324,62 +403,182 @@ async def get_client_credentials(self): # type: ignore[override] ) return NoAuthProvider() - def _instantiate(self, connector_id: str) -> "BaseConnector | None": - connector_cls = get_connector_registry().get(connector_id) - if connector_cls is not None: - cfg = self._configs[connector_id] - auth_provider = self._build_auth_provider(connector_id, cfg.raw) - return connector_cls( - secret_provider=self._secret_provider, - auth_provider=auth_provider, - policy_hook=self._policy_hook, + def _instantiate(self, record: ConfigRecord) -> BaseConnector: + """Build an invoke instance from a resolved config record. I/O-free: + secrets and tokens resolve lazily on first :meth:`BaseConnector.run`.""" + connector_cls = get_connector_registry().get(record.connector_id) + if connector_cls is None: + raise RuntimeError( + f"Connector {record.connector_id!r} has a config but is not registered " + "(filtered by NW_ALLOWED_CONNECTORS or not installed)" ) - raise RuntimeError( - f"Connector {connector_id!r} is enabled in config but not registered " - "(filtered by NW_ALLOWED_CONNECTORS or not installed)" + # Simplified: the __default__ tenant keeps the plain (legacy) secret names + # so existing single-tenant env vars (e.g. EPIC_FHIR_BASE_URL) keep working. + # Named tenants are scoped to NW_{TENANT}_{CONNECTOR}_{CONFIG}_{KEY}. + if record.tenant_id == DEFAULT_TENANT: + scoped: SecretProvider = self._secret_provider + else: + scoped = TenantSecretProvider( + self._secret_provider, + record.tenant_id, + record.connector_id, + config_name=record.name, + ) + + auth_provider = self._build_auth_provider( + record.connector_id, record.raw, secret_provider=scoped + ) + inst = connector_cls( + secret_provider=scoped, + auth_provider=auth_provider, + config=record.raw.get("config", {}), + policy_hook=self._policy_hook, ) + inst._config_name = record.name + return inst + + async def get( + self, + connector_id: str, + *, + tenant_id: str = DEFAULT_TENANT, + config_name: Optional[str] = None, + action: Optional[str] = None, + ) -> BaseConnector: + """Resolve (and cache) the connector instance for a tenant/config. + + Fail-closed: raises :class:`ConfigNotFoundError` when the scope has no + config or ``config_name`` is unknown (bindings map both to 403). Also + callable directly by the embedding application (no binding required). + """ + self._loop = asyncio.get_running_loop() + + # Existence IS entitlement; resolve to the concrete default name first so + # default and explicit callers share one instance. + record = self._store.resolve(tenant_id, connector_id, config_name) + key = (tenant_id, connector_id, record.name) + + with self._instances_guard: + inst = self._instances.get(key) + if inst is not None: + self._instances.move_to_end(key) + return inst + + lock = self._lock_for(key) + async with lock: + with self._instances_guard: + inst = self._instances.get(key) + if inst is not None: + self._instances.move_to_end(key) + return inst + inst = self._instantiate(record) + self._instances[key] = inst + self._evict_if_needed() + return inst + + def is_exposed(self, connector_id: str, protocol: str) -> bool: + """Whether ``connector_id`` may be reached over ``protocol``. + + YAML connectors honour their ``exposed_via`` list; connectors pushed only + via the runtime API (no YAML metadata) are exposed on all protocols. + """ + cfg = self._configs.get(connector_id) + if cfg is None: + return True + return protocol in cfg.exposed_via + + def _lock_for(self, key: Tuple[str, str, str]) -> asyncio.Lock: + with self._locks_guard: + lock = self._locks.get(key) + if lock is None: + lock = asyncio.Lock() + self._locks[key] = lock # locks are never deleted + return lock + + def _evict_if_needed(self) -> None: + while len(self._instances) > self._max_instances: + _, old = self._instances.popitem(last=False) + self._schedule_aclose(old) + + def invalidate_configs(self, tenant_id: str, connector_id: str, names: List[str]) -> None: + """Drop cached instances for the given config names and schedule teardown. + + Called synchronously by the store on every mutating write; may run on a + non-loop thread (store writes are plain sync Python).""" + with self._instances_guard: + for name in names: + old = self._instances.pop((tenant_id, connector_id, name), None) + if old is not None: + self._schedule_aclose(old) + + def _schedule_aclose(self, inst: BaseConnector) -> None: + aclose = getattr(inst, "aclose", None) + if aclose is None: + return # connectors without async cleanup: nothing to do + + async def _safe_aclose() -> None: + try: + await aclose() + except Exception as exc: # noqa: BLE001 + logger.warning("Error during connector aclose: %s", exc) + + try: + loop = asyncio.get_running_loop() + except RuntimeError: + loop = None + + if loop is not None: + loop.create_task(_safe_aclose()) + elif self._loop is not None and self._loop.is_running(): + running_loop = self._loop + running_loop.call_soon_threadsafe(lambda: running_loop.create_task(_safe_aclose())) + # else: no running loop known; the instance is dropped without aclose. + + def _default_instance(self, connector_id: str) -> Optional[BaseConnector]: + """Return (and cache) the ``__default__`` instance for a connector. + + Shares :attr:`_instances` with the async :meth:`get`, so enumeration and + the default-tenant invoke path resolve the SAME object. Sync: + only invoked at import/enumeration and by the default-tenant fast path. + """ + try: + record = self._store.resolve(DEFAULT_TENANT, connector_id, None) + except ConfigNotFoundError: + return None + key = (DEFAULT_TENANT, connector_id, record.name) + with self._instances_guard: + inst = self._instances.get(key) + if inst is None: + inst = self._instantiate(record) + self._instances[key] = inst + else: + self._instances.move_to_end(key) + return inst def get_for_protocol( self, connector_id: str, protocol: str, action: Optional[str] = None ) -> Optional[BaseConnector]: - cfg = self._configs.get(connector_id) - if cfg is None: - logger.warning( - "Requested connector is not configured", - extra={"connector_id": connector_id, "protocol": protocol}, - ) - return None + """Sync accessor for the default-tenant instance (playground/enumeration). - if not cfg.enabled: - logger.warning( - "Requested connector is disabled", - extra={"connector_id": connector_id, "protocol": protocol}, - ) + Returns the SAME instance the async :meth:`get` returns for ``__default__``. + Header-based tenancy still goes through :meth:`get`. + """ + cfg = self._configs.get(connector_id) + if cfg is None or not cfg.enabled: return None - if protocol not in cfg.exposed_via: - logger.warning( - "Connector is not exposed via requested protocol", - extra={"connector_id": connector_id, "protocol": protocol}, - ) return None - - connector = self._connectors.get(connector_id) - if connector is None: - return None - - if action: - logger.debug( - "get_for_protocol resolved connector", - extra={"connector_id": connector_id, "protocol": protocol, "action": action}, - ) - - return connector # type: ignore[return-value] + return self._default_instance(connector_id) def list_for_protocol(self, protocol: str) -> List[BaseConnector]: result: List[BaseConnector] = [] - for connector_id, connector in self._connectors.items(): - if protocol in self._configs[connector_id].exposed_via: - result.append(connector) # type: ignore[arg-type] + for connector_id, cfg in self._configs.items(): + if not cfg.enabled or protocol not in cfg.exposed_via: + continue + if connector_id not in get_connector_registry(): + continue + inst = self._default_instance(connector_id) + if inst is not None: + result.append(inst) return result diff --git a/src/bindings/grpc_server/connector.proto b/src/bindings/grpc_server/connector.proto index f5641910..81f6e58f 100644 --- a/src/bindings/grpc_server/connector.proto +++ b/src/bindings/grpc_server/connector.proto @@ -12,6 +12,7 @@ message InvokeRequest { string connector_id = 1; string action = 2; string payload_json = 3; + optional string config_name = 4; // absent = default config } message InvokeResponse { diff --git a/src/bindings/grpc_server/connector_pb2.py b/src/bindings/grpc_server/connector_pb2.py index 87e0fb61..6d3942bc 100644 --- a/src/bindings/grpc_server/connector_pb2.py +++ b/src/bindings/grpc_server/connector_pb2.py @@ -29,7 +29,7 @@ -DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x0f\x63onnector.proto\x12\x0e\x61ot.connectors\"K\n\rInvokeRequest\x12\x14\n\x0c\x63onnector_id\x18\x01 \x01(\t\x12\x0e\n\x06\x61\x63tion\x18\x02 \x01(\t\x12\x14\n\x0cpayload_json\x18\x03 \x01(\t\"\x83\x01\n\x0eInvokeResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x11\n\tdata_json\x18\x02 \x01(\t\x12\x12\n\nerror_code\x18\x03 \x01(\t\x12\x16\n\x0e\x65rror_category\x18\x04 \x01(\t\x12\x0f\n\x07message\x18\x05 \x01(\t\x12\x10\n\x08trace_id\x18\x06 \x01(\t2[\n\x10\x43onnectorService\x12G\n\x06Invoke\x12\x1d.aot.connectors.InvokeRequest\x1a\x1e.aot.connectors.InvokeResponseb\x06proto3') +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x0f\x63onnector.proto\x12\x0e\x61ot.connectors\"u\n\rInvokeRequest\x12\x14\n\x0c\x63onnector_id\x18\x01 \x01(\t\x12\x0e\n\x06\x61\x63tion\x18\x02 \x01(\t\x12\x14\n\x0cpayload_json\x18\x03 \x01(\t\x12\x18\n\x0b\x63onfig_name\x18\x04 \x01(\tH\x00\x88\x01\x01\x42\x0e\n\x0c_config_name\"\x83\x01\n\x0eInvokeResponse\x12\x0f\n\x07success\x18\x01 \x01(\x08\x12\x11\n\tdata_json\x18\x02 \x01(\t\x12\x12\n\nerror_code\x18\x03 \x01(\t\x12\x16\n\x0e\x65rror_category\x18\x04 \x01(\t\x12\x0f\n\x07message\x18\x05 \x01(\t\x12\x10\n\x08trace_id\x18\x06 \x01(\t2[\n\x10\x43onnectorService\x12G\n\x06Invoke\x12\x1d.aot.connectors.InvokeRequest\x1a\x1e.aot.connectors.InvokeResponseb\x06proto3') _globals = globals() _builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) @@ -37,9 +37,9 @@ if not _descriptor._USE_C_DESCRIPTORS: DESCRIPTOR._loaded_options = None _globals['_INVOKEREQUEST']._serialized_start=35 - _globals['_INVOKEREQUEST']._serialized_end=110 - _globals['_INVOKERESPONSE']._serialized_start=113 - _globals['_INVOKERESPONSE']._serialized_end=244 - _globals['_CONNECTORSERVICE']._serialized_start=246 - _globals['_CONNECTORSERVICE']._serialized_end=337 + _globals['_INVOKEREQUEST']._serialized_end=152 + _globals['_INVOKERESPONSE']._serialized_start=155 + _globals['_INVOKERESPONSE']._serialized_end=286 + _globals['_CONNECTORSERVICE']._serialized_start=288 + _globals['_CONNECTORSERVICE']._serialized_end=379 # @@protoc_insertion_point(module_scope) diff --git a/src/bindings/grpc_server/connector_pb2_grpc.py b/src/bindings/grpc_server/connector_pb2_grpc.py index 64fe11fc..222ef90e 100644 --- a/src/bindings/grpc_server/connector_pb2_grpc.py +++ b/src/bindings/grpc_server/connector_pb2_grpc.py @@ -8,7 +8,7 @@ import grpc import warnings -from . import connector_pb2 as connector__pb2 +import connector_pb2 as connector__pb2 GRPC_GENERATED_VERSION = '1.71.2' GRPC_VERSION = grpc.__version__ diff --git a/src/bindings/grpc_server/server.py b/src/bindings/grpc_server/server.py index 5bb0c8b5..cf8056e2 100644 --- a/src/bindings/grpc_server/server.py +++ b/src/bindings/grpc_server/server.py @@ -17,6 +17,12 @@ from bindings.factory import ConnectorFactory from node_wire_runtime.connector_registry import auto_register from node_wire_runtime import ConnectorResponse, ErrorCategory +from node_wire_runtime.config_store import ConfigNotFoundError +from node_wire_runtime.identity import ( + MissingTenantError, + resolve_config_name, + resolve_tenant_id, +) from node_wire_runtime.ingress import normalize_mcp_tool_arguments from node_wire_runtime.rate_limit import global_rate_limiter, RateLimitExceeded @@ -39,6 +45,7 @@ def __init__(self) -> None: async def _invoke_async( self, request: connector_pb2.InvokeRequest, # type: ignore[name-defined, attr-defined] + metadata: dict[str, str] | None = None, ) -> connector_pb2.InvokeResponse: # type: ignore[name-defined, attr-defined] try: await global_rate_limiter.acquire() @@ -51,8 +58,19 @@ async def _invoke_async( trace_id="", ) - connector = self._factory.get_for_protocol(request.connector_id, "grpc") - if connector is None: + identity = get_grpc_caller_identity() + try: + tenant_id = resolve_tenant_id(headers=metadata or {}, jwt_identity=identity) + except MissingTenantError as exc: + return connector_pb2.InvokeResponse( # type: ignore[name-defined, attr-defined] + success=False, + error_code="MISSING_TENANT", + error_category=ErrorCategory.AUTH.value, + message=str(exc), + trace_id="", + ) + + if not self._factory.is_exposed(request.connector_id, "grpc"): return connector_pb2.InvokeResponse( # type: ignore[name-defined, attr-defined] success=False, error_code="CONNECTOR_NOT_AVAILABLE", @@ -61,6 +79,22 @@ async def _invoke_async( trace_id="", ) + try: + connector = await self._factory.get( + request.connector_id, + tenant_id=tenant_id, + config_name=resolve_config_name(request.config_name or None), + action=request.action, + ) + except ConfigNotFoundError: + return connector_pb2.InvokeResponse( # type: ignore[name-defined, attr-defined] + success=False, + error_code="CONFIG_NOT_FOUND", + error_category=ErrorCategory.AUTH.value, + message="No connector configuration for this tenant", + trace_id="", + ) + payload: Any = {} if request.payload_json: try: @@ -82,11 +116,10 @@ async def _invoke_async( if payload.get("action"): normalize_mcp_tool_arguments(connector, str(payload["action"]), payload) - identity = get_grpc_caller_identity() response: ConnectorResponse = await connector.run( payload, principal=identity.principal if identity else None, - tenant_id=identity.tenant_id if identity else None, + tenant_id=tenant_id, scopes=identity.scopes if identity else None, ) @@ -105,7 +138,10 @@ async def _invoke_async( ) def Invoke(self, request, context): # type: ignore[override] - return _async_runner.run(self._invoke_async(request)) + # gRPC metadata keys are lowercase by spec; tenant header lookup is + # case-insensitive. Pinned once per unary call. + metadata = {k.lower(): v for k, v in (context.invocation_metadata() or ())} + return _async_runner.run(self._invoke_async(request, metadata)) def serve(port: int = 50051) -> None: diff --git a/src/bindings/mcp_server/server.py b/src/bindings/mcp_server/server.py index adff955f..26c7a73d 100644 --- a/src/bindings/mcp_server/server.py +++ b/src/bindings/mcp_server/server.py @@ -20,6 +20,13 @@ log_effective_mcp_auth_state, ) from node_wire_runtime.caller_identity import CallerIdentity +from node_wire_runtime.config_store import ConfigNotFoundError +from node_wire_runtime.identity import ( + MissingTenantError, + is_multitenancy_enabled, + resolve_config_name, + resolve_tenant_id, +) from node_wire_runtime.policies.mcp_scope_policy import ( action_allowed_for_identity_scopes, load_scope_map_from_env, @@ -61,6 +68,13 @@ def is_public_bind_host(host: str) -> bool: default=None, ) +# Pinned factory tenant for the current streamable-http request (§6.3). +# Set with env_pin=None so process NW_TENANT_ID cannot override X-Tenant-ID. +_session_tenant_ctx: ContextVar[str | None] = ContextVar( + "mcp_session_tenant", + default=None, +) + def _process_response_payload(data: Any, max_items: int) -> Tuple[Any, bool, int, Optional[str]]: """ @@ -170,14 +184,18 @@ def __init__( *, server_name: str = "node-wire", connector_ids: Optional[List[str]] = None, + factory: ConnectorFactory | None = None, ) -> None: self._server_name = server_name self._connector_ids: Optional[frozenset[str]] = ( None if connector_ids is None else frozenset(connector_ids) ) auto_register() - self._factory = ConnectorFactory() - self._factory.load() + if factory is not None: + self._factory = factory + else: + self._factory = ConnectorFactory() + self._factory.load() self._upstream_passthrough = _resolve_upstream_passthrough( self._factory, self._connector_ids ) @@ -186,6 +204,8 @@ def __init__( if self._upstream_passthrough else () ) + # §6.4: set in run_stdio from NW_TENANT_ID; unused on HTTP (session pin). + self._stdio_env_tenant_pin: str | None = None try: from importlib.metadata import version as pkg_version @@ -250,11 +270,30 @@ def _list_tools_impl(self, *, identity: CallerIdentity | None = None) -> List[Di f"Manifest contract v{MCP_MANIFEST_CONTRACT_VERSION}." ) ) + # config_name is an optional resolution-time argument (§6.3): an agent + # may target a named config; omitting it uses the tenant's default. + # Only advertise when multitenancy is enabled (ignored when off anyway). + input_schema = entry["input_schema"] + if is_multitenancy_enabled(): + props = input_schema.setdefault("properties", {}) + # Allow null: Groq/OpenAI strict tool validation rejects + # ``config_name: null`` when type is only "string". + props.setdefault( + "config_name", + { + "type": ["string", "null"], + "description": ( + "Optional named connector configuration; omit or null " + "for the tenant default." + ), + }, + ) + tools.append( { "name": f"{cid}.{entry['action']}", "description": tool_desc, - "input_schema": entry["input_schema"], + "input_schema": input_schema, "output_schema": entry["output_schema"], } ) @@ -323,10 +362,54 @@ async def invoke_tool( if self._connector_ids is not None and connector_id not in self._connector_ids: raise ValueError(f"Connector {connector_id!r} is not allowed on this MCP server.") - connector = self._factory.get_for_protocol(connector_id, "mcp") - if connector is None: + if not self._factory.is_exposed(connector_id, "mcp"): raise ValueError(f"Connector {connector_id!r} is not available via MCP.") + # Tenant: HTTP session pin (§6.3) or stdio env pin (§6.4). config_name is a + # resolution-time argument, never a connector input. + arguments = dict(arguments or {}) + config_name = resolve_config_name(arguments.pop("config_name", None)) + # LLMs often fill optional schema keys with null; treat as omitted. + arguments = {k: v for k, v in arguments.items() if v is not None} + session_tenant = _session_tenant_ctx.get() + if session_tenant is not None: + tenant_id = session_tenant + else: + try: + tenant_id = resolve_tenant_id( + headers=_http_request_headers.get(), + jwt_identity=identity, + env_pin=self._stdio_env_tenant_pin, + ) + except MissingTenantError as exc: + raise ValueError(str(exc)) from exc + + try: + connector = await self._factory.get( + connector_id, + tenant_id=tenant_id, + config_name=config_name, + action=action, + ) + except ConfigNotFoundError: + raise ValueError(f"Connector {connector_id!r} is not available via MCP.") + + resolved_config_name = getattr(connector, "_config_name", config_name) + if is_multitenancy_enabled(): + logger.info( + "MCP tool resolved | tool=%s | tenant_id=%s | config_name=%s", + name, + tenant_id, + resolved_config_name or "(default)", + extra={ + "tool_name": name, + "connector_id": connector_id, + "action": action, + "tenant_id": tenant_id, + "config_name": resolved_config_name or "", + }, + ) + run_args = normalize_mcp_tool_arguments(connector, action, arguments) enforce_authoritative_action(run_args, action) run_args["action"] = action @@ -357,7 +440,7 @@ async def invoke_tool( response = await connector.run( run_args, principal=identity.principal if identity else None, - tenant_id=identity.tenant_id if identity else None, + tenant_id=tenant_id, scopes=identity.scopes if identity else None, ) stream_completion_log(trace_id, True, connector_id=connector_id, action=action) @@ -521,6 +604,10 @@ async def _run_stdio_async(self) -> None: from mcp.server.stdio import stdio_server from mcp.server import NotificationOptions + # §6.4: one process, one tenant — pin from env at stdio start (not on HTTP). + raw = os.getenv("NW_TENANT_ID") + self._stdio_env_tenant_pin = raw.strip() if raw and raw.strip() else None + log_effective_mcp_auth_state() low = self._setup_lowlevel_server() @@ -575,10 +662,25 @@ async def dispatch(self, request: Request, call_next): # type: ignore[override] ) setattr(request.state, "nw_mcp_identity", identity) + # §6.3: pin tenant for this HTTP request/session context; never use + # process NW_TENANT_ID (stdio-only) so it cannot override X-Tenant-ID. + try: + session_tenant = resolve_tenant_id( + headers=request.headers, + jwt_identity=identity, + env_pin=None, + ) + except MissingTenantError as exc: + return JSONResponse( + status_code=400, + content={"detail": str(exc), "error_code": "MISSING_TENANT"}, + ) token = _streamable_http_identity_ctx.set(identity) + tenant_token = _session_tenant_ctx.set(session_tenant) try: return await call_next(request) finally: + _session_tenant_ctx.reset(tenant_token) _streamable_http_identity_ctx.reset(token) reset_upstream_passthrough_context() diff --git a/src/bindings/rest_api/app.py b/src/bindings/rest_api/app.py index 15da3abb..759d9e68 100644 --- a/src/bindings/rest_api/app.py +++ b/src/bindings/rest_api/app.py @@ -25,6 +25,18 @@ from node_wire_runtime.connector_registry import auto_register from node_wire_runtime.manifest import build_manifest from node_wire_runtime import ConnectorResponse, ErrorCategory +from node_wire_runtime.config_store import ( + ConfigNameConflictError, + ConfigNotFoundError, + ConfigStoreError, + DefaultDeletionError, +) +from node_wire_runtime.identity import ( + MissingTenantError, + is_multitenancy_enabled, + resolve_config_name, + resolve_tenant_id, +) from node_wire_runtime.ingress import enforce_authoritative_action, normalize_mcp_tool_arguments from opentelemetry import trace from opentelemetry.trace import Status, StatusCode @@ -83,6 +95,14 @@ def _mount_playground(app: FastAPI) -> None: return app.include_router(scenarios_router) + + # Playground feature-flags endpoint: lets the UI query server-side config. + from node_wire_runtime.identity import is_multitenancy_enabled + + @app.get("/playground/feature-flags", tags=["playground"], include_in_schema=False) + async def playground_feature_flags() -> Dict[str, Any]: + return {"multitenancy_enabled": is_multitenancy_enabled()} + app.mount( "/playground", StaticFiles(directory=str(playground_dir), html=True), @@ -111,9 +131,35 @@ def get_factory() -> ConnectorFactory: _factory = ConnectorFactory() auto_register() _factory.load() + from bindings.rest_api.tenant_store import load_tenants + + load_tenants(_factory.store) return _factory +def _persist_tenants(factory_dep: ConnectorFactory) -> None: + from bindings.rest_api.tenant_store import save_tenants + + save_tenants(factory_dep.store) + + +def _pop_secrets(doc: Dict[str, Any]) -> Dict[str, str] | None: + """Extract secrets from a config body. + + Returns ``None`` when the client omitted ``secrets`` (config-only update). + Returns a (possibly empty) dict when ``secrets`` was present — empty means + validate required secrets against existing values for that config only. + """ + if "secrets" not in doc: + return None + raw = doc.pop("secrets") + if raw is None: + return {} + if not isinstance(raw, dict): + raise HTTPException(status_code=400, detail="secrets must be a JSON object") + return {str(k): str(v) for k, v in raw.items() if v is not None and str(v).strip()} + + async def check_rate_limit() -> None: try: # Skip rate limiting if disabled @@ -142,6 +188,294 @@ async def ready() -> Dict[str, str]: return {"status": "ready"} +# --- Runtime connector config API (thin wrappers over ConnectorConfigStore) --- +# Tenant scope comes from the tenant header (never the path). The embedding +# application authenticates callers before these endpoints are reachable +# (RestAuthMiddleware); node-wire does not. + + +def _config_tenant(request: Request) -> str: + try: + return resolve_tenant_id( + headers=request.headers, jwt_identity=get_rest_caller_identity(request) + ) + except MissingTenantError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + + +def _log_tenant_action( + *, + action: str, + tenant_id: str, + connector_id: str | None = None, + config_name: str | None = None, +) -> None: + """Emit tenant context when multitenancy is on (visible in default formatters).""" + if not is_multitenancy_enabled(): + return + # Inline replace so CodeQL treats newline stripping as a sanitizer. + logger.info( + "Tenant config | op=%s | tenant_id=%s | connector=%s | config_name=%s", + str(action).replace("\r", " ").replace("\n", " "), + str(tenant_id).replace("\r", " ").replace("\n", " "), + str(connector_id or "-").replace("\r", " ").replace("\n", " "), + str(config_name or "(default)").replace("\r", " ").replace("\n", " "), + ) + + +def _map_config_error(exc: Exception) -> HTTPException: + if isinstance(exc, MissingTenantError): + return HTTPException(status_code=400, detail=str(exc)) + if isinstance(exc, ConfigNameConflictError): + return HTTPException(status_code=409, detail=str(exc)) + if isinstance(exc, DefaultDeletionError): + return HTTPException(status_code=400, detail=str(exc)) + if isinstance(exc, ConfigNotFoundError): + return HTTPException(status_code=404, detail=str(exc)) + if isinstance(exc, ConfigStoreError): + return HTTPException(status_code=400, detail=str(exc)) + raise exc + + +@app.post("/v1/config/init", tags=["config"]) +async def config_init( + request: Request, + payload: Dict[str, Any], + factory_dep: ConnectorFactory = Depends(get_factory), +) -> JSONResponse: + """Bulk init. With the tenant header: payload is that tenant's + ``{connector_id: [docs]}`` map. Without it: full multi-tenant payload.""" + from node_wire_runtime.identity import tenant_from_headers + + header_tenant = tenant_from_headers(request.headers) + full_payload = {header_tenant: payload} if header_tenant else payload + try: + factory_dep.store.init(full_payload) + _persist_tenants(factory_dep) + except Exception as exc: # noqa: BLE001 + raise _map_config_error(exc) + return JSONResponse(status_code=200, content={"status": "ok"}) + + +@app.get("/v1/tenants", tags=["config"]) +async def list_tenants( + factory_dep: ConnectorFactory = Depends(get_factory), +) -> JSONResponse: + return JSONResponse(status_code=200, content={"tenants": factory_dep.store.list_tenants()}) + + +@app.put("/v1/connectors/{cid}/secrets", tags=["config"]) +async def secrets_upsert( + cid: str, + request: Request, + payload: Dict[str, Any], + factory_dep: ConnectorFactory = Depends(get_factory), +) -> JSONResponse: + from bindings.rest_api.tenant_store import upsert_tenant_secrets + + tenant_id = _config_tenant(request) + secrets = payload.get("secrets") + if not isinstance(secrets, dict): + raise HTTPException(status_code=400, detail="body.secrets must be a JSON object") + config_name = str(payload.get("config_name") or request.query_params.get("config_name") or "").strip() + if not config_name: + raise HTTPException( + status_code=400, + detail="config_name is required (body.config_name or ?config_name=)", + ) + _log_tenant_action( + action="secrets.upsert", + tenant_id=tenant_id, + connector_id=cid, + config_name=config_name, + ) + try: + keys = upsert_tenant_secrets( + tenant_id, + cid, + {str(k): str(v) for k, v in secrets.items()}, + config_name=config_name, + ) + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + factory_dep.invalidate_configs(tenant_id, cid, [config_name]) + _persist_tenants(factory_dep) + return JSONResponse(status_code=200, content={"keys": keys, "config_name": config_name}) + + +@app.get("/v1/connectors/{cid}/secrets", tags=["config"]) +async def secrets_list_keys( + cid: str, + request: Request, + config_name: str = "", +) -> JSONResponse: + from bindings.rest_api.tenant_store import list_secret_logical_keys + + tenant_id = _config_tenant(request) + name = (config_name or request.query_params.get("config_name") or "").strip() + if not name: + raise HTTPException( + status_code=400, + detail="config_name query parameter is required", + ) + return JSONResponse( + status_code=200, + content={"keys": list_secret_logical_keys(tenant_id, cid, name), "config_name": name}, + ) + + +@app.post("/v1/connectors/{cid}/configs", tags=["config"]) +async def config_create( + cid: str, + request: Request, + doc: Dict[str, Any], + factory_dep: ConnectorFactory = Depends(get_factory), +) -> JSONResponse: + from bindings.rest_api.tenant_store import upsert_tenant_secrets + + try: + tenant_id = _config_tenant(request) + body = dict(doc) + secrets = _pop_secrets(body) + config_name = str(body.get("name") or "").strip() + _log_tenant_action( + action="config.create", + tenant_id=tenant_id, + connector_id=cid, + config_name=config_name or None, + ) + if secrets is not None: + try: + upsert_tenant_secrets( + tenant_id, cid, secrets, config_name=config_name + ) + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + try: + record = factory_dep.store.create(tenant_id, cid, body) + except ConfigNameConflictError: + name = str(body.get("name") or "") + record = factory_dep.store.update(tenant_id, cid, name, body) + _persist_tenants(factory_dep) + except HTTPException: + raise + except Exception as exc: # noqa: BLE001 + raise _map_config_error(exc) + return JSONResponse(status_code=201, content={"name": record.name, "default": record.default}) + + +@app.get("/v1/connectors/{cid}/configs", tags=["config"]) +async def config_list( + cid: str, + request: Request, + factory_dep: ConnectorFactory = Depends(get_factory), +) -> JSONResponse: + # No per-list INFO: playground dropdown refresh would flood the terminal. + tenant_id = _config_tenant(request) + return JSONResponse(status_code=200, content=factory_dep.store.list(tenant_id, cid)) + + +@app.get("/v1/connectors/{cid}/configs/{name}", tags=["config"]) +async def config_get( + cid: str, + name: str, + request: Request, + factory_dep: ConnectorFactory = Depends(get_factory), +) -> JSONResponse: + tenant_id = _config_tenant(request) + doc = factory_dep.store.get(tenant_id, cid, name) + if doc is None: + raise HTTPException(status_code=404, detail="Config not found") + return JSONResponse(status_code=200, content=doc) + + +@app.put("/v1/connectors/{cid}/configs/{name}", tags=["config"]) +async def config_update( + cid: str, + name: str, + request: Request, + doc: Dict[str, Any], + factory_dep: ConnectorFactory = Depends(get_factory), +) -> JSONResponse: + from bindings.rest_api.tenant_store import upsert_tenant_secrets + + try: + tenant_id = _config_tenant(request) + body = dict(doc) + secrets = _pop_secrets(body) + _log_tenant_action( + action="config.update", + tenant_id=tenant_id, + connector_id=cid, + config_name=name, + ) + if secrets: + try: + upsert_tenant_secrets(tenant_id, cid, secrets, config_name=name) + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + record = factory_dep.store.update(tenant_id, cid, name, body) + _persist_tenants(factory_dep) + except HTTPException: + raise + except Exception as exc: # noqa: BLE001 + raise _map_config_error(exc) + return JSONResponse(status_code=200, content={"name": record.name, "default": record.default}) + + +@app.delete("/v1/connectors/{cid}/configs/{name}", tags=["config"]) +async def config_delete( + cid: str, + name: str, + request: Request, + new_default: str | None = None, + factory_dep: ConnectorFactory = Depends(get_factory), +) -> JSONResponse: + try: + tenant_id = _config_tenant(request) + _log_tenant_action( + action="config.delete", + tenant_id=tenant_id, + connector_id=cid, + config_name=name, + ) + factory_dep.store.delete(tenant_id, cid, name, new_default=new_default) + # Secrets are per named config; clear this config's vault. + from bindings.rest_api.tenant_store import ( + clear_config_secrets, + clear_tenant_connector_secrets, + ) + + clear_config_secrets(tenant_id, cid, name) + if not factory_dep.store.has_config(tenant_id, cid): + clear_tenant_connector_secrets(tenant_id, cid) + _persist_tenants(factory_dep) + except Exception as exc: # noqa: BLE001 + raise _map_config_error(exc) + return JSONResponse(status_code=200, content={"status": "ok"}) + + +@app.put("/v1/connectors/{cid}/configs/{name}/default", tags=["config"]) +async def config_set_default( + cid: str, + name: str, + request: Request, + factory_dep: ConnectorFactory = Depends(get_factory), +) -> JSONResponse: + try: + tenant_id = _config_tenant(request) + _log_tenant_action( + action="config.set_default", + tenant_id=tenant_id, + connector_id=cid, + config_name=name, + ) + factory_dep.store.set_default(tenant_id, cid, name) + _persist_tenants(factory_dep) + except Exception as exc: # noqa: BLE001 + raise _map_config_error(exc) + return JSONResponse(status_code=200, content={"status": "ok"}) + def _http_status_for_category(category: ErrorCategory | None) -> int: if category is None: return 200 @@ -196,10 +530,28 @@ async def endpoint( span = trace.get_current_span() span.set_attribute("connector.id", cid) span.set_attribute("connector.action", act) + + rest_id = get_rest_caller_identity(request) + try: + tenant_id = resolve_tenant_id(headers=request.headers, jwt_identity=rest_id) + except MissingTenantError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + + run_payload = dict(payload) + # config_name is a resolution-time argument, never a connector input. + # Suppressed when multitenancy is disabled so legacy path is always used. + config_name = resolve_config_name(run_payload.pop("config_name", None)) + _log_tenant_action( + action=f"rest.{act}", + tenant_id=tenant_id, + connector_id=cid, + config_name=config_name, + ) + if _rate_limit_enabled(): limiter = _get_rate_limiter() identity_key = get_request_identity_key(request) - rate_key = f"{cid}:{act}:{identity_key}" + rate_key = f"{tenant_id}:{cid}:{act}:{identity_key}" result = limiter.consume(rate_key) if not result.allowed: return JSONResponse( @@ -208,10 +560,26 @@ async def endpoint( headers={"Retry-After": str(result.retry_after_seconds)}, ) - connector = factory_dep.get_for_protocol(cid, "rest", action=act) - if connector is None: + if not factory_dep.is_exposed(cid, "rest"): raise HTTPException(status_code=404, detail="Connector not available for REST") - run_payload = dict(payload) + + try: + connector = await factory_dep.get( + cid, tenant_id=tenant_id, config_name=config_name, action=act + ) + except ConfigNotFoundError: + # Unknown scope and unknown config name return the same body so config + # names cannot be enumerated (fail-closed). + raise HTTPException( + status_code=403, + detail=( + "No connector configuration for this tenant. " + "Use Add config on this connector page." + ), + ) + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + run_payload = normalize_mcp_tool_arguments(connector, act, run_payload) try: enforce_authoritative_action(run_payload, act) @@ -220,11 +588,10 @@ async def endpoint( run_payload["action"] = act # Let the runtime (Layer A) perform full schema validation. # Any validation errors will be mapped into ConnectorResponse. - rest_id = get_rest_caller_identity(request) response: ConnectorResponse = await connector.run( run_payload, principal=rest_id.principal if rest_id else None, - tenant_id=rest_id.tenant_id if rest_id else None, + tenant_id=tenant_id, scopes=rest_id.scopes if rest_id else None, ) status = _http_status_for_category(response.error_category) diff --git a/src/bindings/rest_api/tenant_store.py b/src/bindings/rest_api/tenant_store.py new file mode 100644 index 00000000..95582229 --- /dev/null +++ b/src/bindings/rest_api/tenant_store.py @@ -0,0 +1,485 @@ +# +# SPDX-FileCopyrightText: 2026 AOT Technologies +# SPDX-License-Identifier: Apache-2.0 +# +"""Multi-tenant config + secret overlay persistence for Node Wire REST. + +Simplified: one YAML file rewritten on each mutation; gitignored by the repo. +""" + +from __future__ import annotations + +import json +import logging +import os +import re +import threading +from pathlib import Path +from typing import Any, Callable, Dict, List, Mapping, Tuple + +import yaml + +from node_wire_runtime.config_store import ConfigNotFoundError, ConnectorConfigStore +from node_wire_runtime.secrets import OverlaySecretProvider, tenant_scoped_secret_key + +logger = logging.getLogger("bindings.rest_api.tenant_store") + +_REPO_ROOT = Path(__file__).resolve().parent.parent.parent.parent +DEFAULT_TENANTS_PATH = _REPO_ROOT / "config" / "tenants.yaml" +_LEGACY_TENANTS_PATH = _REPO_ROOT / "config" / "playground_tenants.yaml" + +# (process env name, logical secret key used by connector / auth refs) +SHARED_ENV_BY_CONNECTOR: Dict[str, List[Tuple[str, str]]] = { + "fhir_epic": [ + ("EPIC_FHIR_BASE_URL", "epic_fhir_base_url"), + ("EPIC_TOKEN_URL", "EPIC_TOKEN_URL"), + ], + "fhir_cerner": [ + ("CERNER_FHIR_BASE_URL", "cerner_fhir_base_url"), + ("CERNER_TOKEN_URL", "CERNER_TOKEN_URL"), + ("CERNER_TOKEN_URL", "cerner_token_url"), + ], + "salesforce": [ + ("SALESFORCE_TOKEN_URL", "SALESFORCE_TOKEN_URL"), + ("SALESFORCE_INSTANCE_URL", "salesforce_instance_url"), + ], +} + +# When a secrets map is provided, these logical keys must be present and non-empty. +REQUIRED_SECRETS_BY_CONNECTOR: Dict[str, List[str]] = { + "google_drive": ["GOOGLE_DRIVE_SA_JSON"], + "fhir_epic": ["EPIC_CLIENT_ID", "EPIC_PRIVATE_KEY", "EPIC_KID"], + "fhir_cerner": ["CERNER_CLIENT_ID", "CERNER_PRIVATE_KEY", "CERNER_KID"], + "slack": ["SLACK_BOT_TOKEN"], + "stripe": ["stripe_api_key"], + "salesforce": [ + "SALESFORCE_CLIENT_ID", + "SALESFORCE_CLIENT_SECRET", + "SALESFORCE_REFRESH_TOKEN", + ], +} + +# Format kind per logical secret key (validated only for newly supplied values). +SECRET_FORMAT_BY_CONNECTOR: Dict[str, Dict[str, str]] = { + "google_drive": {"GOOGLE_DRIVE_SA_JSON": "google_sa_json"}, + "fhir_epic": { + "EPIC_CLIENT_ID": "opaque_secret", + "EPIC_PRIVATE_KEY": "pem_private_key", + "EPIC_KID": "jwt_kid", + }, + "fhir_cerner": { + "CERNER_CLIENT_ID": "opaque_secret", + "CERNER_PRIVATE_KEY": "pem_private_key", + "CERNER_KID": "jwt_kid", + "CERNER_SCOPES": "scopes_space_separated", + }, + "slack": {"SLACK_BOT_TOKEN": "slack_bot_token"}, + "stripe": {"stripe_api_key": "stripe_secret_key"}, + "salesforce": { + "SALESFORCE_CLIENT_ID": "opaque_secret", + "SALESFORCE_CLIENT_SECRET": "opaque_secret", + "SALESFORCE_REFRESH_TOKEN": "opaque_secret", + }, +} + +_PEM_PRIVATE_RE = re.compile( + r"-----BEGIN (?:RSA |EC |OPENSSH )?PRIVATE KEY-----[\s\S]+?-----END (?:RSA |EC |OPENSSH )?PRIVATE KEY-----" +) +_JWT_KID_RE = re.compile(r"^[A-Za-z0-9._\-]{1,128}$") +_SLACK_BOT_RE = re.compile(r"^xoxb-[A-Za-z0-9\-]+$") +_STRIPE_SECRET_RE = re.compile(r"^sk_(?:test|live)_[A-Za-z0-9]+$") + + +def _normalize_pem(value: str) -> str: + return value.replace("\\n", "\n").strip() + + +def _validate_pem_private_key(value: str) -> None: + pem = _normalize_pem(value) + if not _PEM_PRIVATE_RE.search(pem): + if re.search(r"-----BEGIN (?:RSA )?PUBLIC KEY-----", pem): + raise ValueError("expected a PEM private key, not a public key") + if "BEGIN CERTIFICATE" in pem: + raise ValueError("expected a PEM private key, not a certificate") + raise ValueError("expected PEM private key (BEGIN/END PRIVATE KEY block)") + + +def _validate_jwt_kid(value: str) -> None: + if not _JWT_KID_RE.match(value.strip()): + raise ValueError("expected kid as 1–128 chars of A–Z, a–z, 0–9, '.', '_', or '-'") + + +def _validate_google_sa_json(value: str) -> None: + try: + data = json.loads(value) + except json.JSONDecodeError as exc: + raise ValueError(f"expected JSON service account: {exc.msg}") from exc + if not isinstance(data, dict): + raise ValueError("expected a JSON object for service account") + if data.get("type") != "service_account": + raise ValueError('expected JSON with "type": "service_account"') + + +def _validate_slack_bot_token(value: str) -> None: + if not _SLACK_BOT_RE.match(value.strip()): + raise ValueError("expected Slack bot token starting with xoxb-") + + +def _validate_stripe_secret_key(value: str) -> None: + if not _STRIPE_SECRET_RE.match(value.strip()): + raise ValueError("expected Stripe secret key (sk_test_… or sk_live_…)") + + +def _validate_opaque_secret(value: str) -> None: + v = value.strip() + if not v: + raise ValueError("expected a non-empty value") + if "\n" in v or "\r" in v: + raise ValueError("must be a single-line value") + + +def _validate_scopes_space_separated(value: str) -> None: + parts = value.split() + if not parts: + raise ValueError("expected one or more scopes separated by spaces") + if any(not p for p in parts): + raise ValueError("scopes must be non-empty tokens separated by spaces") + + +_FORMAT_VALIDATORS: Dict[str, Callable[[str], None]] = { + "pem_private_key": _validate_pem_private_key, + "jwt_kid": _validate_jwt_kid, + "google_sa_json": _validate_google_sa_json, + "slack_bot_token": _validate_slack_bot_token, + "stripe_secret_key": _validate_stripe_secret_key, + "opaque_secret": _validate_opaque_secret, + "scopes_space_separated": _validate_scopes_space_separated, +} + + +def validate_required_secrets(connector_id: str, logical_secrets: Mapping[str, str]) -> None: + """Raise ValueError when required varying keys are missing from an effective secrets map.""" + required = REQUIRED_SECRETS_BY_CONNECTOR.get(connector_id) or [] + missing = [ + key + for key in required + if key not in logical_secrets or not str(logical_secrets.get(key) or "").strip() + ] + if missing: + raise ValueError(f"missing required secrets for {connector_id}: {', '.join(missing)}") + + +def validate_secret_formats(connector_id: str, logical_secrets: Mapping[str, str]) -> None: + """Raise ValueError when supplied secret values fail format checks for this connector.""" + formats = SECRET_FORMAT_BY_CONNECTOR.get(connector_id) or {} + errors: List[str] = [] + for key, value in logical_secrets.items(): + fmt = formats.get(str(key)) + if not fmt: + continue + raw = str(value).strip() + if not raw: + continue + validator = _FORMAT_VALIDATORS.get(fmt) + if not validator: + continue + try: + validator(raw) + except ValueError as exc: + errors.append(f"{key}: {exc}") + if errors: + raise ValueError("; ".join(errors)) + + +_lock = threading.RLock() +# Faithful nested secrets: tenant → connector → config_name → logical_key → value. +_nested_secrets_mirror: Dict[str, Dict[str, Dict[str, Dict[str, str]]]] = {} + + +def existing_logical_secrets( + tenant_id: str, connector_id: str, config_name: str +) -> Dict[str, str]: + return dict( + ((_nested_secrets_mirror.get(tenant_id) or {}).get(connector_id) or {}).get( + config_name + ) + or {} + ) + + +def tenants_path(*, for_write: bool = False) -> Path: + """Resolve the tenants YAML path. + + ``NW_TENANTS_PATH`` wins; otherwise write to ``config/tenants.yaml``. + Reads may fall back to legacy ``config/playground_tenants.yaml`` if present. + """ + override = ( + os.environ.get("NW_TENANTS_PATH", "").strip() + or os.environ.get("NW_PLAYGROUND_TENANTS_PATH", "").strip() + ) + if override: + return Path(override) + if for_write or DEFAULT_TENANTS_PATH.is_file() or not _LEGACY_TENANTS_PATH.is_file(): + return DEFAULT_TENANTS_PATH + return _LEGACY_TENANTS_PATH + + +def _export_store(store: ConnectorConfigStore) -> Dict[str, Any]: + out: Dict[str, Any] = {} + with store._lock: # noqa: SLF001 — same-process persist + for tenant_id, connectors in store._data.items(): # noqa: SLF001 + cid_map: Dict[str, List[Dict[str, Any]]] = {} + for connector_id, records in connectors.items(): + cid_map[connector_id] = [dict(rec.raw) for rec in records.values()] + out[tenant_id] = cid_map + return out + + +def _export_secrets_mirror() -> Dict[str, Any]: + return { + t: {c: {cfg: dict(kv) for cfg, kv in configs.items()} for c, configs in cons.items()} + for t, cons in _nested_secrets_mirror.items() + } + + +def save_tenants(store: ConnectorConfigStore) -> None: + path = tenants_path(for_write=True) + with _lock: + payload = { + "tenants": _export_store(store), + "secrets": _export_secrets_mirror(), + } + path.parent.mkdir(parents=True, exist_ok=True) + tmp = path.with_suffix(".yaml.tmp") + with open(tmp, "w", encoding="utf-8") as f: + yaml.safe_dump( + payload, f, default_flow_style=False, allow_unicode=True, sort_keys=False + ) + tmp.replace(path) + logger.info("Wrote tenants file", extra={"path": str(path)}) + + +def upsert_tenant_secrets( + tenant_id: str, + connector_id: str, + logical_secrets: Mapping[str, str], + *, + config_name: str, + auto_shared_env: bool = True, + require_varying: bool = True, +) -> List[str]: + """Merge logical secrets for one named config. Returns logical keys set for that config. + + Empty / omitted keys keep existing values for this config only (partial update). + New configs must supply required secrets — sibling configs are not shared. + """ + name = (config_name or "").strip() + if not name: + raise ValueError("config_name is required for tenant secrets") + + overlay = OverlaySecretProvider.instance() + merged: Dict[str, str] = { + str(k).strip(): str(v) + for k, v in logical_secrets.items() + if k is not None and str(k).strip() and v is not None and str(v).strip() + } + existing = existing_logical_secrets(tenant_id, connector_id, name) + if auto_shared_env: + for env_key, logical_key in SHARED_ENV_BY_CONNECTOR.get(connector_id, []): + if logical_key in merged and merged[logical_key].strip(): + continue + if logical_key in existing: + continue + host_val = os.environ.get(env_key) + if host_val is None: + host_val = os.environ.get(env_key.lower()) + if host_val is not None and str(host_val).strip(): + merged[logical_key] = str(host_val).strip() + + if require_varying: + effective = {**existing, **merged} + validate_required_secrets(connector_id, effective) + + if merged: + validate_secret_formats(connector_id, merged) + + flat: Dict[str, str] = {} + for logical_key, value in merged.items(): + scoped = tenant_scoped_secret_key( + tenant_id, connector_id, logical_key, config_name=name + ) + flat[scoped] = value + _nested_secrets_mirror.setdefault(tenant_id, {}).setdefault(connector_id, {}).setdefault( + name, {} + )[logical_key] = value + + if flat: + overlay.set_many(flat) + return sorted( + (_nested_secrets_mirror.get(tenant_id) or {}) + .get(connector_id, {}) + .get(name, {}) + .keys() + ) + + +def _is_legacy_connector_secrets(kv: Mapping[str, Any]) -> bool: + """True when secrets are flat logical→string (pre per-config nesting).""" + if not kv: + return False + return all(not isinstance(v, dict) for v in kv.values()) + + +def load_tenants(store: ConnectorConfigStore) -> None: + path = tenants_path() + if not path.is_file(): + return + with _lock: + with open(path, encoding="utf-8") as f: + raw = yaml.safe_load(f) or {} + if not isinstance(raw, dict): + logger.warning("Ignoring invalid tenants file", extra={"path": str(path)}) + return + + tenants = raw.get("tenants") or {} + if isinstance(tenants, dict): + for tenant_id, connectors in tenants.items(): + if not isinstance(connectors, dict): + continue + for connector_id, docs in connectors.items(): + if not isinstance(docs, list): + continue + for doc in docs: + if not isinstance(doc, dict) or not doc.get("name"): + continue + name = str(doc["name"]) + try: + store.create(tenant_id, connector_id, doc) + except Exception: + try: + store.update(tenant_id, connector_id, name, doc) + except ConfigNotFoundError: + logger.warning( + "Could not load tenant config", + extra={ + "tenant_id": tenant_id, + "connector_id": connector_id, + "name": name, + }, + ) + + global _nested_secrets_mirror + _nested_secrets_mirror = {} + flat: Dict[str, str] = {} + secrets = raw.get("secrets") or {} + if isinstance(secrets, dict): + for tenant_id, connectors in secrets.items(): + if not isinstance(connectors, dict): + continue + for connector_id, kv in connectors.items(): + if not isinstance(kv, dict): + continue + if _is_legacy_connector_secrets(kv): + # Recommended: no auto-migrate — re-enter per-config credentials. + logger.warning( + "Skipping legacy connector-scoped secrets; " + "re-save credentials per named config", + extra={ + "tenant_id": tenant_id, + "connector_id": connector_id, + }, + ) + continue + for config_name, logical_map in kv.items(): + if not isinstance(logical_map, dict): + continue + cfg = str(config_name) + for logical, value in logical_map.items(): + scoped = tenant_scoped_secret_key( + tenant_id, + connector_id, + str(logical), + config_name=cfg, + ) + flat[scoped] = str(value) + _nested_secrets_mirror.setdefault(tenant_id, {}).setdefault( + connector_id, {} + ).setdefault(cfg, {})[str(logical)] = str(value) + OverlaySecretProvider.instance().replace_all(flat) + logger.info( + "Loaded tenants file", + extra={ + "path": str(path), + "tenants": len(tenants) if isinstance(tenants, dict) else 0, + }, + ) + + +def list_secret_logical_keys( + tenant_id: str, connector_id: str, config_name: str +) -> List[str]: + name = (config_name or "").strip() + if not name: + return [] + return sorted( + (_nested_secrets_mirror.get(tenant_id) or {}) + .get(connector_id, {}) + .get(name, {}) + .keys() + ) + + +def clear_config_secrets(tenant_id: str, connector_id: str, config_name: str) -> None: + """Drop overlay + mirror secrets for one named config.""" + name = (config_name or "").strip() + if not name: + return + with _lock: + cons = _nested_secrets_mirror.get(tenant_id) or {} + configs = cons.get(connector_id) or {} + logical_map = configs.pop(name, {}) + if connector_id in cons and not cons[connector_id]: + del cons[connector_id] + if tenant_id in _nested_secrets_mirror and not _nested_secrets_mirror[tenant_id]: + del _nested_secrets_mirror[tenant_id] + + overlay = OverlaySecretProvider.instance() + data = overlay.export() + keys_to_drop = set(logical_map.keys()) | set( + overlay.logical_keys_for(tenant_id, connector_id, config_name=name) + ) + for logical_key in keys_to_drop: + data.pop( + tenant_scoped_secret_key( + tenant_id, connector_id, logical_key, config_name=name + ), + None, + ) + overlay.replace_all(data) + + +def clear_tenant_connector_secrets(tenant_id: str, connector_id: str) -> None: + """Drop overlay + mirror secrets for all configs under one tenant/connector.""" + with _lock: + configs = (_nested_secrets_mirror.get(tenant_id) or {}).pop(connector_id, {}) + if tenant_id in _nested_secrets_mirror and not _nested_secrets_mirror[tenant_id]: + del _nested_secrets_mirror[tenant_id] + + overlay = OverlaySecretProvider.instance() + data = overlay.export() + for config_name, logical_map in configs.items(): + for logical_key in set(logical_map.keys()) | set( + overlay.logical_keys_for( + tenant_id, connector_id, config_name=config_name + ) + ): + data.pop( + tenant_scoped_secret_key( + tenant_id, + connector_id, + logical_key, + config_name=config_name, + ), + None, + ) + overlay.replace_all(data) diff --git a/src/node_wire_runtime/__init__.py b/src/node_wire_runtime/__init__.py index 58c09f06..e0b8ac94 100644 --- a/src/node_wire_runtime/__init__.py +++ b/src/node_wire_runtime/__init__.py @@ -4,9 +4,25 @@ # from .models import ConnectorResponse, ErrorCategory from .errors import ErrorMapper -from .secrets import SecretProvider, EnvSecretProvider, SecretNotFoundError, SecretProviderError +from .secrets import ( + SecretProvider, + EnvSecretProvider, + SecretNotFoundError, + SecretProviderError, + TenantSecretNotFoundError, + TenantSecretProvider, +) from .policy import PolicyHook, PolicyDenied from .caller_identity import CallerIdentity, build_caller_identity +from .config_store import ( + ConnectorConfigStore, + ConfigRecord, + ConfigNotFoundError, + ConfigNameConflictError, + DefaultDeletionError, + DEFAULT_TENANT, +) +from .identity import resolve_tenant_id, tenant_from_headers from .auth import ( AuthProvider, NoAuthProvider, @@ -44,10 +60,20 @@ "EnvSecretProvider", "SecretNotFoundError", "SecretProviderError", + "TenantSecretNotFoundError", + "TenantSecretProvider", "PolicyHook", "PolicyDenied", "CallerIdentity", "build_caller_identity", + "ConnectorConfigStore", + "ConfigRecord", + "ConfigNotFoundError", + "ConfigNameConflictError", + "DefaultDeletionError", + "DEFAULT_TENANT", + "resolve_tenant_id", + "tenant_from_headers", "AuthProvider", "NoAuthProvider", "StaticTokenAuthProvider", diff --git a/src/node_wire_runtime/base_connector.py b/src/node_wire_runtime/base_connector.py index e217c40b..9b1f7bdf 100644 --- a/src/node_wire_runtime/base_connector.py +++ b/src/node_wire_runtime/base_connector.py @@ -378,12 +378,19 @@ def __init__( secret_provider: Optional[SecretProvider] = None, policy_hook: Optional[PolicyHook] = None, auth_provider: Optional[AuthProvider] = None, + config: Optional[Dict[str, Any]] = None, ) -> None: cls = type(self) self._input_model_cls = cls._union_input_model self._output_model_cls = cls.output_model self._secret_provider = secret_provider self._policy_hook = policy_hook + # Per-config settings (e.g. channel, base_url) from the resolved config + # record. Available for connectors that opt in; existing connectors keep + # reading settings from secrets. The concrete config name is set by the + # factory as ``_config_name`` for observability. + self.config: Dict[str, Any] = config or {} + self._config_name: Optional[str] = None # Default to NoAuthProvider (null-object) so connectors never receive None. self._auth_provider: AuthProvider = ( auth_provider if auth_provider is not None else NoAuthProvider() @@ -454,33 +461,40 @@ async def run( - Maps exceptions into the standard error taxonomy """ trace_id = str(uuid.uuid4()) - _start = time.monotonic() + config_name = getattr(self, "_config_name", None) with tracer.start_as_current_span( "connector.run", attributes={ "connector.id": self.connector_id, "connector.action": self.action, + "config.name": config_name or "", "tenant.id": tenant_id or "", "principal.id": principal or "", "trace.id": trace_id, }, ): logger.info( - "Starting connector execution", + "Starting connector execution | connector=%s | action=%s | tenant_id=%s | config_name=%s", + self.connector_id, + self.action, + tenant_id or "(none)", + config_name or "(default)", extra={ "trace_id": trace_id, "connector_id": self.connector_id, + "config_name": config_name, "action": self.action, "principal": principal, - "tenant_id": tenant_id, "scopes": list(scopes) if scopes else [], "audit": True, "audit_event": "invocation_start", + "tenant_id": tenant_id or "", }, ) _response: Optional[ConnectorResponse] = None + _start = time.monotonic() token = _caller_execution_ctx.set((principal, tenant_id, scopes)) try: try: diff --git a/src/node_wire_runtime/config_store.py b/src/node_wire_runtime/config_store.py new file mode 100644 index 00000000..f5362128 --- /dev/null +++ b/src/node_wire_runtime/config_store.py @@ -0,0 +1,406 @@ +# +# SPDX-FileCopyrightText: 2026 AOT Technologies +# SPDX-License-Identifier: Apache-2.0 +# +""" +node_wire_runtime.config_store +============================== + +In-memory connector configuration store for header-based multi-tenancy. + +The embedding application pushes connector configuration at runtime (there is no +file/DB/Vault read at request time). Each ``(tenant_id, connector_id)`` scope may +hold any number of uniquely named configs; exactly one is the default. + +The store is plain in-process memory and thread-safe. Every mutating call +synchronously invalidates the affected factory instances via a factory attached +with :meth:`ConnectorConfigStore.attach_factory` (same process, so no TTL and no +change-detection polling). + +Redaction: inline secret-bearing fields (suffix markers below) are masked on every +read surface (:meth:`get`, :meth:`list`). Only the factory's internal +:meth:`resolve` sees the unredacted record. +""" + +from __future__ import annotations + +import copy +import threading +from dataclasses import dataclass, field +from typing import Any, Dict, List, Optional + +DEFAULT_TENANT = "__default__" + +# Inline secret fields end with one of these suffixes; reference fields end with +# ``_secret`` / ``_key`` and carry bare logical names (never redacted). +SECRET_FIELD_MARKERS = ( + "_value", + "_secret_value", + "password", + "token_value", + "private_key_value", +) + +_REDACTION_MASK = "\u2022" * 8 # eight bullet characters + + +class ConfigStoreError(Exception): + """Base class for config store errors.""" + + +class ConfigNotFoundError(ConfigStoreError): + """No config exists for the scope, or the named config is unknown. + + Deliberately raised for both cases so config names cannot be enumerated + (unknown scope and unknown name are indistinguishable to the caller). + """ + + +class ConfigNameConflictError(ConfigStoreError): + """A config with the same name already exists in the scope.""" + + +class DefaultDeletionError(ConfigStoreError): + """Deleting the current default requires nominating ``new_default``.""" + + +@dataclass +class ConfigRecord: + """A single resolved config for one ``(tenant, connector)`` scope. + + ``raw`` is the full config document (``config`` / ``auth`` blocks). ``exposed_via`` + carries protocol gating from the YAML bootstrap; it is internal and not part of + the public config-document shape. + """ + + tenant_id: str + connector_id: str + name: str + default: bool + raw: Dict[str, Any] + exposed_via: List[str] = field(default_factory=list) + + +def redact(doc: Any) -> Any: + """Return a deep copy of ``doc`` with inline secret values masked. + + Reference fields (bare logical names) pass through unchanged. Applied to + every read surface: GET, list responses, logs, error messages. + """ + if isinstance(doc, dict): + out: Dict[str, Any] = {} + for k, v in doc.items(): + if isinstance(v, (dict, list)): + out[k] = redact(v) + elif isinstance(k, str) and any(k.endswith(m) for m in SECRET_FIELD_MARKERS): + out[k] = _REDACTION_MASK + else: + out[k] = v + return out + if isinstance(doc, list): + return [redact(item) for item in doc] + return doc + + +def _validate_doc(doc: Any) -> Dict[str, Any]: + """Validate a single config document at the trust boundary. + + Callers (bindings) may pass arbitrary JSON; enforce shape here rather than + failing obscurely deeper in the factory. + """ + if not isinstance(doc, dict): + raise ConfigStoreError("config document must be a JSON object") + name = doc.get("name") + if not isinstance(name, str) or not name.strip(): + raise ConfigStoreError("config document requires a non-empty string 'name'") + if "default" in doc and not isinstance(doc["default"], bool): + raise ConfigStoreError("config 'default' must be a boolean") + if "config" in doc and not isinstance(doc["config"], dict): + raise ConfigStoreError("config 'config' block must be a JSON object") + if "auth" in doc and not isinstance(doc["auth"], dict): + raise ConfigStoreError("config 'auth' block must be a JSON object") + return doc + + +class ConnectorConfigStore: + """In-memory store. Thread-safe. Every write synchronously invalidates the + affected factory instances (same process, so no TTL and no polling).""" + + def __init__(self) -> None: + # tenant_id -> connector_id -> name -> ConfigRecord (insertion-ordered) + self._data: Dict[str, Dict[str, Dict[str, ConfigRecord]]] = {} + self._lock = threading.RLock() + self._factory: Any = None + + def attach_factory(self, factory: Any) -> None: + """Attach the factory that receives synchronous invalidation on writes.""" + self._factory = factory + + # ---- write path ----------------------------------------------------- + + def init(self, payload: Dict[str, Any]) -> None: + """Bulk load: replaces ALL configs. + + Shape: ``{ tenant_id: { connector_id: [ config_doc, ... ] } }``. Exactly + one doc per ``(tenant, connector)`` may set ``default=true``; if none does, + the first in the list becomes default. Invalidates every previously cached + instance. + """ + if not isinstance(payload, dict): + raise ConfigStoreError("init payload must be a JSON object") + + new_data: Dict[str, Dict[str, Dict[str, ConfigRecord]]] = {} + for tenant_id, connectors in payload.items(): + if not isinstance(connectors, dict): + raise ConfigStoreError( + f"init payload for tenant {tenant_id!r} must be a JSON object" + ) + tenant_map: Dict[str, Dict[str, ConfigRecord]] = {} + for connector_id, docs in connectors.items(): + if not isinstance(docs, list): + raise ConfigStoreError( + f"init payload for {tenant_id!r}/{connector_id!r} must be a list" + ) + tenant_map[connector_id] = self._build_scope(tenant_id, connector_id, docs) + new_data[tenant_id] = tenant_map + + with self._lock: + # Collect the full previous key set so init() invalidates everything. + previous: List[tuple[str, str, str]] = [] + for tenant_id, connectors in self._data.items(): + for connector_id, records in connectors.items(): + for name in records: + previous.append((tenant_id, connector_id, name)) + self._data = new_data + for tenant_id, connector_id, name in previous: + self._invalidate(tenant_id, connector_id, [name]) + + def _build_scope( + self, tenant_id: str, connector_id: str, docs: List[Any] + ) -> Dict[str, ConfigRecord]: + """Build the name -> record map for one scope, enforcing default rules.""" + records: Dict[str, ConfigRecord] = {} + default_seen: Optional[str] = None + for doc in docs: + doc = _validate_doc(doc) + name = doc["name"] + if name in records: + raise ConfigNameConflictError( + f"duplicate config name {name!r} for {tenant_id!r}/{connector_id!r}" + ) + is_default = bool(doc.get("default", False)) + if is_default and default_seen is not None: + raise ConfigStoreError( + f"more than one default config for {tenant_id!r}/{connector_id!r} " + f"({default_seen!r} and {name!r})" + ) + if is_default: + default_seen = name + records[name] = ConfigRecord( + tenant_id=tenant_id, + connector_id=connector_id, + name=name, + default=is_default, + raw=copy.deepcopy(doc), + exposed_via=list(doc.get("exposed_via", []) or []), + ) + if records and default_seen is None: + # First config becomes the default when none is explicitly marked. + first_name = next(iter(records)) + records[first_name].default = True + return records + + def create(self, tenant_id: str, connector_id: str, doc: Dict[str, Any]) -> ConfigRecord: + """Add one named config. Name must be unique in scope. The first config + for a scope becomes the default automatically.""" + doc = _validate_doc(doc) + name = doc["name"] + with self._lock: + scope = self._data.setdefault(tenant_id, {}).setdefault(connector_id, {}) + if name in scope: + raise ConfigNameConflictError( + f"config {name!r} already exists for {tenant_id!r}/{connector_id!r}" + ) + is_first = len(scope) == 0 + requested_default = bool(doc.get("default", False)) + make_default = is_first or requested_default + record = ConfigRecord( + tenant_id=tenant_id, + connector_id=connector_id, + name=name, + default=make_default, + raw=copy.deepcopy(doc), + exposed_via=list(doc.get("exposed_via", []) or []), + ) + if make_default: + for other in scope.values(): + other.default = False + scope[name] = record + self._invalidate(tenant_id, connector_id, [name]) + return record + + def update( + self, tenant_id: str, connector_id: str, name: str, doc: Dict[str, Any] + ) -> ConfigRecord: + """Replace the named config's contents. The name itself is immutable + (delete + create to rename), keeping instance keys unambiguous.""" + doc = _validate_doc(doc) + if doc["name"] != name: + raise ConfigStoreError( + f"config name is immutable: cannot rename {name!r} to {doc['name']!r} " + "(delete + create instead)" + ) + with self._lock: + record = self._require(tenant_id, connector_id, name) + record.raw = copy.deepcopy(doc) + record.exposed_via = list(doc.get("exposed_via", []) or []) + # 'default' flag is managed via set_default/delete, not update. + self._invalidate(tenant_id, connector_id, [name]) + return record + + def set_default(self, tenant_id: str, connector_id: str, name: str) -> None: + """Make ``name`` the sole default for the scope.""" + with self._lock: + scope = self._scope(tenant_id, connector_id) + if scope is None or name not in scope: + raise ConfigNotFoundError(f"no config {name!r} for {tenant_id!r}/{connector_id!r}") + changed: List[str] = [] + for other_name, other in scope.items(): + should_be = other_name == name + if other.default != should_be: + other.default = should_be + changed.append(other_name) + # Moving the default reroutes which concrete key default calls resolve + # to; invalidate the previous default so cached default routing refreshes. + if changed: + self._invalidate(tenant_id, connector_id, changed) + + def delete( + self, + tenant_id: str, + connector_id: str, + name: str, + new_default: Optional[str] = None, + ) -> None: + """Delete the named config. + + Deleting the current default requires ``new_default`` unless it is the last + config, in which case the connector is removed for the tenant entirely + (subsequent invokes fail closed with :class:`ConfigNotFoundError`). + """ + with self._lock: + scope = self._scope(tenant_id, connector_id) + if scope is None or name not in scope: + raise ConfigNotFoundError(f"no config {name!r} for {tenant_id!r}/{connector_id!r}") + record = scope[name] + is_last = len(scope) == 1 + + if record.default and not is_last: + if new_default is None: + raise DefaultDeletionError( + f"config {name!r} is the default for {tenant_id!r}/{connector_id!r}; " + "provide new_default to nominate a replacement" + ) + if new_default == name or new_default not in scope: + raise ConfigNotFoundError( + f"new_default {new_default!r} is not a config for " + f"{tenant_id!r}/{connector_id!r}" + ) + + del scope[name] + invalidated = [name] + if not scope: + # Last config removed: drop the scope so has_config() -> False. + del self._data[tenant_id][connector_id] + if not self._data[tenant_id]: + del self._data[tenant_id] + elif record.default: + scope[new_default].default = True # type: ignore[index] + invalidated.append(new_default) # type: ignore[arg-type] + self._invalidate(tenant_id, connector_id, invalidated) + + # ---- read path (redacted) ------------------------------------------ + + def get(self, tenant_id: str, connector_id: str, name: str) -> Optional[Dict[str, Any]]: + """Read one config with secret-bearing fields redacted. ``None`` if absent.""" + with self._lock: + scope = self._scope(tenant_id, connector_id) + if scope is None or name not in scope: + return None + return self._public_view(scope[name]) + + def list(self, tenant_id: str, connector_id: Optional[str] = None) -> List[Dict[str, Any]]: + """Redacted list. ``connector_id=None`` lists all configs for the tenant.""" + with self._lock: + tenant_map = self._data.get(tenant_id) + if tenant_map is None: + return [] + out: List[Dict[str, Any]] = [] + connector_ids = [connector_id] if connector_id is not None else list(tenant_map) + for cid in connector_ids: + scope = tenant_map.get(cid) + if not scope: + continue + for record in scope.values(): + out.append(self._public_view(record)) + return out + + def has_config(self, tenant_id: str, connector_id: str) -> bool: + """True when at least one config exists for the scope (entitlement check).""" + with self._lock: + scope = self._scope(tenant_id, connector_id) + return bool(scope) + + def list_tenants(self) -> List[str]: + """Tenant ids that currently have at least one connector config.""" + with self._lock: + return sorted(self._data.keys()) + + # ---- internal (unredacted) ----------------------------------------- + + def resolve(self, tenant_id: str, connector_id: str, name: Optional[str]) -> ConfigRecord: + """INTERNAL (factory use). ``name=None`` resolves the default. Missing scope + or name raises :class:`ConfigNotFoundError`. Returns the UNREDACTED record.""" + with self._lock: + scope = self._scope(tenant_id, connector_id) + if not scope: + raise ConfigNotFoundError( + f"no config for tenant {tenant_id!r} / connector {connector_id!r}" + ) + if name is None: + for record in scope.values(): + if record.default: + return record + # Defensive: a non-empty scope always has a default by construction. + return next(iter(scope.values())) + found = scope.get(name) + if found is None: + raise ConfigNotFoundError( + f"no config for tenant {tenant_id!r} / connector {connector_id!r}" + ) + return found + + # ---- helpers -------------------------------------------------------- + + def _scope(self, tenant_id: str, connector_id: str) -> Optional[Dict[str, ConfigRecord]]: + tenant_map = self._data.get(tenant_id) + if tenant_map is None: + return None + return tenant_map.get(connector_id) + + def _require(self, tenant_id: str, connector_id: str, name: str) -> ConfigRecord: + scope = self._scope(tenant_id, connector_id) + if scope is None or name not in scope: + raise ConfigNotFoundError(f"no config {name!r} for {tenant_id!r}/{connector_id!r}") + return scope[name] + + @staticmethod + def _public_view(record: ConfigRecord) -> Dict[str, Any]: + view = redact(record.raw) + view["name"] = record.name + view["default"] = record.default + return view + + def _invalidate(self, tenant_id: str, connector_id: str, names: List[str]) -> None: + if self._factory is not None: + self._factory.invalidate_configs(tenant_id, connector_id, names) diff --git a/src/node_wire_runtime/identity.py b/src/node_wire_runtime/identity.py new file mode 100644 index 00000000..23b84985 --- /dev/null +++ b/src/node_wire_runtime/identity.py @@ -0,0 +1,129 @@ +# +# SPDX-FileCopyrightText: 2026 AOT Technologies +# SPDX-License-Identifier: Apache-2.0 +# +""" +node_wire_runtime.identity +========================== + +Header/argument-based tenant identity for the embedding application. + +Node Wire is a library embedded in a trusted host. The tenant id it supplies is +taken at face value (the ``X-Tenant-ID`` header is fully trusted, no verification). +If bindings are exposed to untrusted clients, deploy an authenticating gateway in +front; node-wire itself performs no authentication here. + +Resolution order when NW_MULTITENANCY_ENABLED=true (see plan decision C / C2): + 1. ``env_pin`` transport with no headers (MCP stdio ``NW_TENANT_ID``) + 2. header ``X-Tenant-ID`` / ``NW_TENANT_ID_HEADER`` (case-insensitive) + 3. ``jwt_identity`` existing JWT-claim tenancy (unchanged) + 4. raise MissingTenantError when none of the above are present + +When NW_MULTITENANCY_ENABLED=false (default), all calls return DEFAULT_TENANT +regardless of headers, JWT, or env_pin so connectors behave exactly as before +multi-tenancy was introduced. +""" + +from __future__ import annotations + +import os +from typing import Any, Mapping, Optional + +from node_wire_runtime.config_store import DEFAULT_TENANT + +# Header name lookup is case-insensitive; store the configured name lowercased. +TENANT_HEADER = os.getenv("NW_TENANT_ID_HEADER", "x-tenant-id").lower() + +MISSING_TENANT_MESSAGE = "X-Tenant-ID is required when multitenancy is enabled" + + +class MissingTenantError(ValueError): + """Raised when multitenancy is enabled and no tenant id can be resolved.""" + + def __init__(self, message: str = MISSING_TENANT_MESSAGE) -> None: + super().__init__(message) + + +def is_multitenancy_enabled() -> bool: + """Return True when NW_MULTITENANCY_ENABLED is set to a truthy value. + + Defaults to False so existing single-tenant deployments are unaffected. + """ + return os.getenv("NW_MULTITENANCY_ENABLED", "false").strip().lower() in ( + "1", + "true", + "yes", + "on", + ) + + +def tenant_from_headers(headers: Optional[Mapping[str, str]]) -> Optional[str]: + """Return the tenant id from a case-insensitive header lookup, or ``None``. + + An empty / whitespace-only value is treated as absent. + """ + if not headers: + return None + for k, v in headers.items(): + if isinstance(k, str) and k.lower() == TENANT_HEADER: + if v is None: + return None + stripped = v.strip() + return stripped or None + return None + + +def resolve_tenant_id( + *, + headers: Optional[Mapping[str, str]] = None, + jwt_identity: Any = None, + env_pin: Optional[str] = None, +) -> str: + """Resolve the effective tenant id. + + When NW_MULTITENANCY_ENABLED is false (default), always returns DEFAULT_TENANT + regardless of inputs so connectors behave as legacy single-tenant. + + When enabled, requires env_pin, header, or jwt tenant claim; otherwise raises + :class:`MissingTenantError`. Explicit ``__default__`` in the header is allowed. + + ``jwt_identity`` is any object exposing a ``tenant_id`` attribute (e.g. + :class:`~node_wire_runtime.caller_identity.CallerIdentity`). + """ + # Simplified: early exit keeps all callers unchanged; no second code path. + if not is_multitenancy_enabled(): + return DEFAULT_TENANT + + if env_pin is not None: + pinned = env_pin.strip() + if pinned: + return pinned + + header_tenant = tenant_from_headers(headers) + if header_tenant: + return header_tenant + + if jwt_identity is not None: + claim = getattr(jwt_identity, "tenant_id", None) + if claim is not None and str(claim).strip(): + return str(claim).strip() + + raise MissingTenantError() + + +def resolve_config_name(config_name: Optional[str]) -> Optional[str]: + """Return a non-empty config name when multitenancy is enabled, else None. + + Ensures user-supplied named configs are silently ignored in single-tenant + mode so the factory falls back to the YAML-bootstrapped default. + + ``None``, non-strings, and blank strings are treated as omit (tenant default). + LLMs often emit ``config_name: null`` for optional fields; that must not + fail closed as an unknown name. + """ + if not is_multitenancy_enabled(): + return None + if not isinstance(config_name, str): + return None + stripped = config_name.strip() + return stripped or None diff --git a/src/node_wire_runtime/mcp_client/client.py b/src/node_wire_runtime/mcp_client/client.py index 44a88aa3..da43f024 100644 --- a/src/node_wire_runtime/mcp_client/client.py +++ b/src/node_wire_runtime/mcp_client/client.py @@ -43,6 +43,7 @@ def __init__( token_manager: Optional[TokenManager] = None, http_client: Optional[httpx.AsyncClient] = None, reauthorize: Optional[Callable[[], Awaitable[OAuthTokenSet]]] = None, + extra_headers: Optional[Dict[str, str]] = None, ) -> None: self._base_url = base_url.rstrip("/") self._config = config or config_from_env(server_url=base_url) @@ -51,6 +52,7 @@ def __init__( self._initialized = False self._http = http_client self._owns_http = http_client is None + self._extra_headers = dict(extra_headers or {}) self._token_manager = token_manager or TokenManager( self._config, user_id=self._user_id, @@ -85,6 +87,7 @@ async def _auth_headers(self) -> Dict[str, str]: def _merge_headers(self, extra: Dict[str, str]) -> Dict[str, str]: out = dict(extra) + out.update(self._extra_headers) if self._session_id: out["Mcp-Session-Id"] = self._session_id return out @@ -204,6 +207,7 @@ def create_http_mcp_client( user_id: Optional[str] = None, force_oauth: bool = False, reauthorize: Optional[Callable[[], Awaitable[OAuthTokenSet]]] = None, + extra_headers: Optional[Dict[str, str]] = None, ): """ Factory: OAuth client when enabled, else legacy static-token HTTP client. @@ -216,12 +220,17 @@ def create_http_mcp_client( from agents.toolhive import ToolHiveMcpClient if legacy_static_mcp_token() and not force_oauth: - return ToolHiveMcpClient(base_url) + return ToolHiveMcpClient(base_url, extra_headers=extra_headers) if mcp_oauth_enabled() or force_oauth: - return McpOAuthClient(base_url, user_id=user_id, reauthorize=reauthorize) + return McpOAuthClient( + base_url, + user_id=user_id, + reauthorize=reauthorize, + extra_headers=extra_headers, + ) - return ToolHiveMcpClient(base_url) + return ToolHiveMcpClient(base_url, extra_headers=extra_headers) def create_http_mcp_clients_for_urls( @@ -229,5 +238,14 @@ def create_http_mcp_clients_for_urls( *, user_id: Optional[str] = None, reauthorize: Optional[Callable[[], Awaitable[OAuthTokenSet]]] = None, + extra_headers: Optional[Dict[str, str]] = None, ) -> list: - return [create_http_mcp_client(u, user_id=user_id, reauthorize=reauthorize) for u in urls] + return [ + create_http_mcp_client( + u, + user_id=user_id, + reauthorize=reauthorize, + extra_headers=extra_headers, + ) + for u in urls + ] diff --git a/src/node_wire_runtime/policy.py b/src/node_wire_runtime/policy.py index ab45f7dc..b1ccd8e8 100644 --- a/src/node_wire_runtime/policy.py +++ b/src/node_wire_runtime/policy.py @@ -6,7 +6,10 @@ from abc import ABC, abstractmethod from dataclasses import dataclass -from typing import Any, Mapping, Optional +from typing import TYPE_CHECKING, Any, Mapping, Optional + +if TYPE_CHECKING: + from node_wire_runtime.config_store import ConnectorConfigStore @dataclass @@ -39,3 +42,21 @@ def check(self, context: PolicyContext) -> None: Raise PolicyDenied with a human-readable message when execution is not allowed. """ raise NotImplementedError + + +class TenantConfigHook(PolicyHook): + """Defense in depth for ``configured = entitled``. + + The factory already fails closed when a scope has no config; this hook catches + any code path that reached :meth:`BaseConnector.run` without factory resolution. + """ + + def __init__(self, store: "ConnectorConfigStore") -> None: + self._store = store + + def check(self, context: PolicyContext) -> None: + tenant_id = context.tenant_id or "__default__" + if not self._store.has_config(tenant_id, context.connector_id): + raise PolicyDenied( + f"no config for tenant '{tenant_id}' / connector '{context.connector_id}'" + ) diff --git a/src/node_wire_runtime/secrets/__init__.py b/src/node_wire_runtime/secrets/__init__.py index beb097c4..f41a3a3e 100644 --- a/src/node_wire_runtime/secrets/__init__.py +++ b/src/node_wire_runtime/secrets/__init__.py @@ -24,16 +24,24 @@ from node_wire_runtime.secrets.base import ( EnvSecretProvider, + OverlaySecretProvider, SecretNotFoundError, SecretProvider, SecretProviderError, + TenantSecretNotFoundError, + TenantSecretProvider, + tenant_scoped_secret_key, ) from node_wire_runtime.secrets.chained import ChainedSecretProvider __all__ = [ "SecretProvider", "EnvSecretProvider", + "OverlaySecretProvider", "SecretNotFoundError", "SecretProviderError", + "TenantSecretNotFoundError", + "TenantSecretProvider", + "tenant_scoped_secret_key", "ChainedSecretProvider", ] diff --git a/src/node_wire_runtime/secrets/base.py b/src/node_wire_runtime/secrets/base.py index 337dbd0f..c72b87df 100644 --- a/src/node_wire_runtime/secrets/base.py +++ b/src/node_wire_runtime/secrets/base.py @@ -12,6 +12,10 @@ class SecretNotFoundError(KeyError): """The requested key does not exist in this provider.""" +class TenantSecretNotFoundError(SecretNotFoundError): + """A tenant-scoped secret is absent. Strict: never falls back to a shared value.""" + + class SecretProviderError(RuntimeError): """The provider itself failed (auth, network, config). Do not swallow.""" @@ -61,3 +65,146 @@ def get_secret(self, key: str) -> str: if self._legacy_empty_on_missing: return "" raise SecretNotFoundError(key) + + +def _sanitize_secret_segment(segment: str) -> str: + """Uppercase and replace every non-alphanumeric char with ``_`` for env names.""" + return "".join(ch if ch.isalnum() else "_" for ch in segment).upper() + + +def tenant_scoped_secret_key( + tenant_id: str, + connector_id: str, + logical_key: str, + *, + config_name: str | None = None, +) -> str: + """Build ``NW_{TENANT}_{CONNECTOR}_{KEY}`` or ``NW_{TENANT}_{CONNECTOR}_{CONFIG}_{KEY}``. + + When ``config_name`` is set (named multitenant configs), secrets are isolated + per config. Omit ``config_name`` for legacy tenant+connector scoping. + """ + parts = [ + "NW", + _sanitize_secret_segment(tenant_id), + _sanitize_secret_segment(connector_id), + ] + if config_name is not None and str(config_name).strip(): + parts.append(_sanitize_secret_segment(str(config_name).strip())) + parts.append(_sanitize_secret_segment(logical_key)) + return "_".join(parts) + + +class OverlaySecretProvider(SecretProvider): + """In-memory secret map checked before env (tenant credentials overlay). + + Keys match :func:`tenant_scoped_secret_key` ( + ``NW_{TENANT}_{CONNECTOR}_{KEY}`` or ``NW_{TENANT}_{CONNECTOR}_{CONFIG}_{KEY}``). + Process-wide singleton via :meth:`instance`. + """ + + _instance: "OverlaySecretProvider | None" = None + + def __init__(self) -> None: + self._data: dict[str, str] = {} + + @classmethod + def instance(cls) -> "OverlaySecretProvider": + if cls._instance is None: + cls._instance = cls() + return cls._instance + + def get_secret(self, key: str) -> str: + val = self._data.get(key) + if val is not None: + return val + val = self._data.get(key.upper()) + if val is not None: + return val + raise SecretNotFoundError(key) + + def set_secret(self, key: str, value: str) -> None: + self._data[key] = value + + def set_many(self, mapping: dict[str, str]) -> None: + for k, v in mapping.items(): + self._data[str(k)] = str(v) + + def clear(self) -> None: + self._data.clear() + + def replace_all(self, mapping: dict[str, str]) -> None: + self._data = {str(k): str(v) for k, v in mapping.items()} + + def export(self) -> dict[str, str]: + return dict(self._data) + + def logical_keys_for( + self, + tenant_id: str, + connector_id: str, + *, + config_name: str | None = None, + ) -> list[str]: + """Return logical key names present for this tenant/connector[/config] scope.""" + parts = [ + "NW", + _sanitize_secret_segment(tenant_id), + _sanitize_secret_segment(connector_id), + ] + if config_name is not None and str(config_name).strip(): + parts.append(_sanitize_secret_segment(str(config_name).strip())) + prefix = "_".join(parts) + "_" + out: list[str] = [] + for full in self._data: + if not full.startswith(prefix): + continue + logical = full[len(prefix) :] + if logical: + out.append(logical) + return sorted(out) + + +class TenantSecretProvider(SecretProvider): + """Scopes secret lookups to ``{tenant}/{connector}[/{config}]/{key}``. + + Delegates to an inner :class:`SecretProvider`, translating the logical path to + ``NW_{TENANT}_{CONNECTOR}_{KEY}`` or ``NW_{TENANT}_{CONNECTOR}_{CONFIG}_{KEY}``. + Strict: a missing secret raises :class:`TenantSecretNotFoundError`. + + ``key`` is the bare logical name carried by a config's reference field + (e.g. ``GOOGLE_DRIVE_SA_JSON``). + """ + + def __init__( + self, + inner: SecretProvider, + tenant_id: str, + connector_id: str, + *, + config_name: str | None = None, + ) -> None: + self._inner = inner + self._tenant_id = tenant_id + self._connector_id = connector_id + self._config_name = (config_name or "").strip() or None + + def _scoped_key(self, key: str) -> str: + return tenant_scoped_secret_key( + self._tenant_id, + self._connector_id, + key, + config_name=self._config_name, + ) + + def get_secret(self, key: str) -> str: + scoped = self._scoped_key(key) + try: + return self._inner.get_secret(scoped) + except SecretNotFoundError as exc: + scope = f"{self._tenant_id}/{self._connector_id}" + if self._config_name: + scope = f"{scope}/{self._config_name}" + raise TenantSecretNotFoundError( + f"tenant secret not found: {scope}/{key} (resolved key {scoped!r})" + ) from exc diff --git a/tests/conftest.py b/tests/conftest.py index 9d97f679..f74a095b 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -69,6 +69,9 @@ def _preload_connector_logic_modules() -> None: def _rest_auth_disabled_for_tests(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv("NW_REST_AUTH_DISABLED", "true") monkeypatch.setenv("NW_MCP_AUTH_DISABLED", "true") + # Isolate from developer .env: legacy REST tests expect single-tenant unless + # a test explicitly enables multitenancy. + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "false") monkeypatch.delenv("GOOGLE_DRIVE_AUTH_PROVIDER", raising=False) monkeypatch.setenv("NW_MCP_SCOPE_POLICY_DEFAULT", "allow") monkeypatch.setenv("NW_JWT_AUDIENCE", "node-wire-test") @@ -76,3 +79,18 @@ def _rest_auth_disabled_for_tests(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv("NW_RATE_LIMIT_BURST", "1000") # Increase for tests monkeypatch.setenv("NW_RATE_LIMIT_REFILL_RATE", "100.0") # Increase for tests monkeypatch.setenv("NW_RATE_LIMIT_DISABLED", "true") # Disable rate limiting for tests + + +@pytest.fixture(autouse=True) +def _tenants_isolated(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + """Keep tenants.yaml off the repo and clear overlay between tests.""" + path = tmp_path / "tenants.yaml" + monkeypatch.setenv("NW_TENANTS_PATH", str(path)) + from node_wire_runtime.secrets import OverlaySecretProvider + import bindings.rest_api.tenant_store as pt + + OverlaySecretProvider.instance().clear() + pt._nested_secrets_mirror.clear() + yield + OverlaySecretProvider.instance().clear() + pt._nested_secrets_mirror.clear() diff --git a/tests/test_factory_and_rest.py b/tests/test_factory_and_rest.py index 8c83c8e7..aa8b0387 100644 --- a/tests/test_factory_and_rest.py +++ b/tests/test_factory_and_rest.py @@ -16,7 +16,7 @@ def test_factory_loads_config(monkeypatch: pytest.MonkeyPatch): - monkeypatch.setattr(ConnectorFactory, "_instantiate", lambda self, cid: MagicMock()) + monkeypatch.setattr(ConnectorFactory, "_instantiate", lambda self, record: MagicMock()) factory = ConnectorFactory() factory.load() @@ -111,9 +111,8 @@ def test_rest_post_with_bearer_succeeds_when_key_required(monkeypatch: pytest.Mo monkeypatch.setenv("NW_REST_API_KEY", "unit-test-secret") monkeypatch.setenv("NW_RATE_LIMIT_DISABLED", "true") # Disable rate limiting for this test - mock_factory = MagicMock() - mock_factory.get_for_protocol.return_value = _stub_connector( - ConnectorResponse(success=True, data={"ok": True}, trace_id="t-rest") + mock_factory = _mock_factory( + _stub_connector(ConnectorResponse(success=True, data={"ok": True}, trace_id="t-rest")) ) app.dependency_overrides[get_factory] = lambda: mock_factory try: @@ -137,10 +136,8 @@ def test_rest_post_propagates_api_key_identity_to_connector_run( monkeypatch.delenv("NW_REST_API_KEY_SCOPES", raising=False) monkeypatch.setenv("NW_REST_API_KEY", "unit-test-secret") - mock_factory = MagicMock() - mock_factory.get_for_protocol.return_value = _stub_connector( - ConnectorResponse(success=True, data={}, trace_id="t-p") - ) + stub = _stub_connector(ConnectorResponse(success=True, data={}, trace_id="t-p")) + mock_factory = _mock_factory(stub) app.dependency_overrides[get_factory] = lambda: mock_factory try: client = TestClient(app) @@ -152,16 +149,18 @@ def test_rest_post_propagates_api_key_identity_to_connector_run( finally: app.dependency_overrides.clear() - stub = mock_factory.get_for_protocol.return_value kwargs = stub.run.await_args.kwargs assert kwargs["principal"] == "api-key-user" - assert kwargs["tenant_id"] is None + # No tenant header and no JWT claim -> normalized to the default sentinel. + assert kwargs["tenant_id"] == "__default__" assert kwargs["scopes"] == () def test_rest_post_propagates_jwt_claims_to_connector_run(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.delenv("NW_REST_AUTH_DISABLED", raising=False) monkeypatch.delenv("NW_REST_API_KEY", raising=False) + # JWT tenant claim is only applied when multitenancy is enabled. + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") secret = "rest-jwt-test-secret-at-least-32bytes!!" monkeypatch.setenv("NW_REST_JWT_SECRET", secret) @@ -170,10 +169,8 @@ def test_rest_post_propagates_jwt_claims_to_connector_run(monkeypatch: pytest.Mo secret, ) - mock_factory = MagicMock() - mock_factory.get_for_protocol.return_value = _stub_connector( - ConnectorResponse(success=True, data={}, trace_id="t-j") - ) + stub = _stub_connector(ConnectorResponse(success=True, data={}, trace_id="t-j")) + mock_factory = _mock_factory(stub) app.dependency_overrides[get_factory] = lambda: mock_factory try: client = TestClient(app) @@ -186,7 +183,6 @@ def test_rest_post_propagates_jwt_claims_to_connector_run(monkeypatch: pytest.Mo app.dependency_overrides.clear() # The connector needs to be called first to set up the mock - stub = mock_factory.get_for_protocol.return_value assert stub.run is not None, "Connector mock was not called" kwargs = stub.run.await_args.kwargs assert kwargs["principal"] == "alice" @@ -221,11 +217,19 @@ def _stub_connector(response: ConnectorResponse) -> MagicMock: return c +def _mock_factory(stub: MagicMock) -> MagicMock: + """Factory mock for the async, tenant-aware invoke path (``get`` + ``is_exposed``).""" + f = MagicMock() + f.is_exposed.return_value = True + f.get = AsyncMock(return_value=stub) + return f + + def test_rest_post_connector_success() -> None: """Dynamic POST forwards payload to connector.run and returns JSON with 200.""" resp_body = ConnectorResponse(success=True, data={"ok": True}, trace_id="t-rest") - mock_factory = MagicMock() - mock_factory.get_for_protocol.return_value = _stub_connector(resp_body) + stub = _stub_connector(resp_body) + mock_factory = _mock_factory(stub) app.dependency_overrides[get_factory] = lambda: mock_factory try: @@ -240,8 +244,9 @@ def test_rest_post_connector_success() -> None: body = r.json() assert body["success"] is True assert body["trace_id"] == "t-rest" - mock_factory.get_for_protocol.assert_called_with("http_generic", "rest", action="request") - stub = mock_factory.get_for_protocol.return_value + mock_factory.get.assert_awaited_with( + "http_generic", tenant_id="__default__", config_name=None, action="request" + ) stub.run.assert_awaited_once() call_payload = stub.run.await_args[0][0] assert call_payload["action"] == "request" @@ -250,9 +255,8 @@ def test_rest_post_connector_success() -> None: def test_rest_post_connector_rejects_conflicting_action_in_body() -> None: """Body action must match URL path segment (same as MCP tool name authority).""" - mock_factory = MagicMock() - mock_factory.get_for_protocol.return_value = _stub_connector( - ConnectorResponse(success=True, data={}, trace_id="t") + mock_factory = _mock_factory( + _stub_connector(ConnectorResponse(success=True, data={}, trace_id="t")) ) app.dependency_overrides[get_factory] = lambda: mock_factory @@ -288,8 +292,7 @@ def test_rest_post_connector_error_category_http_status( error_code="E1", message="nope", ) - mock_factory = MagicMock() - mock_factory.get_for_protocol.return_value = _stub_connector(resp_body) + mock_factory = _mock_factory(_stub_connector(resp_body)) app.dependency_overrides[get_factory] = lambda: mock_factory try: @@ -307,7 +310,7 @@ def test_rest_post_connector_error_category_http_status( def test_rest_post_connector_not_available_returns_404() -> None: mock_factory = MagicMock() - mock_factory.get_for_protocol.return_value = None + mock_factory.is_exposed.return_value = False app.dependency_overrides[get_factory] = lambda: mock_factory try: @@ -362,3 +365,70 @@ def test_factory_default_scope_policy_is_deny_without_env( factory = ConnectorFactory() assert factory._policy_hook is not None + + +@pytest.mark.asyncio +async def test_factory_instances_cache_safe_under_concurrent_get_and_invalidate( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """_instances is shared by async get, sync invalidate, and _default_instance.""" + import asyncio + import threading + + from node_wire_runtime.config_store import DEFAULT_TENANT + + factory = ConnectorFactory() + factory.store.init( + { + DEFAULT_TENANT: { + "http_generic": [{"name": "default", "default": True, "config": {}}], + } + } + ) + + def _fake_instantiate(self: ConnectorFactory, record: object) -> MagicMock: + m = MagicMock() + m.aclose = AsyncMock() + return m + + monkeypatch.setattr(ConnectorFactory, "_instantiate", _fake_instantiate) + + errors: list[Exception] = [] + + async def hammer_get() -> None: + try: + for _ in range(40): + await factory.get("http_generic", tenant_id=DEFAULT_TENANT) + except Exception as exc: # noqa: BLE001 + errors.append(exc) + + def hammer_invalidate() -> None: + try: + for _ in range(40): + factory.invalidate_configs(DEFAULT_TENANT, "http_generic", ["default"]) + except Exception as exc: # noqa: BLE001 + errors.append(exc) + + def hammer_default() -> None: + try: + for _ in range(40): + factory._default_instance("http_generic") + except Exception as exc: # noqa: BLE001 + errors.append(exc) + + threads = [ + threading.Thread(target=hammer_invalidate), + threading.Thread(target=hammer_default), + threading.Thread(target=hammer_invalidate), + ] + for t in threads: + t.start() + await asyncio.gather(hammer_get(), hammer_get(), hammer_get()) + for t in threads: + t.join() + + assert errors == [] + # Cache remains a usable OrderedDict after concurrent mutation. + with factory._instances_guard: + _ = list(factory._instances.items()) + assert len(factory._instances) <= 1 diff --git a/tests/test_grpc_bindings.py b/tests/test_grpc_bindings.py index d239ea27..147ef59c 100644 --- a/tests/test_grpc_bindings.py +++ b/tests/test_grpc_bindings.py @@ -201,7 +201,7 @@ async def _raise(*_a: Any, **_kw: Any) -> None: async def test_invoke_unknown_connector_returns_not_available( servicer: ConnectorServiceServicer, ) -> None: - with patch.object(servicer._factory, "get_for_protocol", return_value=None): + with patch.object(servicer._factory, "is_exposed", return_value=False): req = connector_pb2.InvokeRequest(connector_id="no_such", action="act") resp = await servicer._invoke_async(req) @@ -213,7 +213,10 @@ async def test_invoke_unknown_connector_returns_not_available( async def test_invoke_invalid_json_payload(servicer: ConnectorServiceServicer) -> None: fake_connector = MagicMock() - with patch.object(servicer._factory, "get_for_protocol", return_value=fake_connector): + with ( + patch.object(servicer._factory, "is_exposed", return_value=True), + patch.object(servicer._factory, "get", new=AsyncMock(return_value=fake_connector)), + ): req = connector_pb2.InvokeRequest( connector_id="x", action="y", payload_json="not-valid-json{" ) @@ -233,7 +236,10 @@ async def test_invoke_success_path(servicer: ConnectorServiceServicer) -> None: trace_id="trace-001", ) ) - with patch.object(servicer._factory, "get_for_protocol", return_value=fake_connector): + with ( + patch.object(servicer._factory, "is_exposed", return_value=True), + patch.object(servicer._factory, "get", new=AsyncMock(return_value=fake_connector)), + ): req = connector_pb2.InvokeRequest( connector_id="x", action="greet", payload_json='{"field": "val"}' ) @@ -253,7 +259,10 @@ async def mock_run(payload: Any, **_: Any) -> ConnectorResponse: fake_connector = MagicMock() fake_connector.run = mock_run - with patch.object(servicer._factory, "get_for_protocol", return_value=fake_connector): + with ( + patch.object(servicer._factory, "is_exposed", return_value=True), + patch.object(servicer._factory, "get", new=AsyncMock(return_value=fake_connector)), + ): req = connector_pb2.InvokeRequest( connector_id="x", action="do_thing", @@ -266,6 +275,7 @@ async def mock_run(payload: Any, **_: Any) -> ConnectorResponse: async def test_invoke_identity_propagated(servicer: ConnectorServiceServicer) -> None: from node_wire_runtime.caller_identity import build_caller_identity + from node_wire_runtime.config_store import DEFAULT_TENANT identity = build_caller_identity({"sub": "grpc-svc"}, auth_type="grpc_api_key") captured: list[Any] = [] @@ -277,17 +287,21 @@ async def mock_run(payload: Any, **kwargs: Any) -> ConnectorResponse: fake_connector = MagicMock() fake_connector.run = mock_run with ( - patch.object(servicer._factory, "get_for_protocol", return_value=fake_connector), + patch.object(servicer._factory, "is_exposed", return_value=True), + patch.object(servicer._factory, "get", new=AsyncMock(return_value=fake_connector)), patch("bindings.grpc_server.server.get_grpc_caller_identity", return_value=identity), ): req = connector_pb2.InvokeRequest(connector_id="x", action="act", payload_json="{}") await servicer._invoke_async(req) assert captured[0]["principal"] == identity.principal - assert captured[0]["tenant_id"] == identity.tenant_id + # Multitenancy off in tests: resolve_tenant_id always returns DEFAULT_TENANT. + assert captured[0]["tenant_id"] == DEFAULT_TENANT async def test_invoke_no_identity_passes_none(servicer: ConnectorServiceServicer) -> None: + from node_wire_runtime.config_store import DEFAULT_TENANT + captured: list[Any] = [] async def mock_run(payload: Any, **kwargs: Any) -> ConnectorResponse: @@ -297,20 +311,24 @@ async def mock_run(payload: Any, **kwargs: Any) -> ConnectorResponse: fake_connector = MagicMock() fake_connector.run = mock_run with ( - patch.object(servicer._factory, "get_for_protocol", return_value=fake_connector), + patch.object(servicer._factory, "is_exposed", return_value=True), + patch.object(servicer._factory, "get", new=AsyncMock(return_value=fake_connector)), patch("bindings.grpc_server.server.get_grpc_caller_identity", return_value=None), ): req = connector_pb2.InvokeRequest(connector_id="x", action="act", payload_json="{}") await servicer._invoke_async(req) assert captured[0]["principal"] is None - assert captured[0]["tenant_id"] is None + assert captured[0]["tenant_id"] == DEFAULT_TENANT async def test_invoke_empty_payload_json(servicer: ConnectorServiceServicer) -> None: fake_connector = MagicMock() fake_connector.run = AsyncMock(return_value=ConnectorResponse(success=True, trace_id="t4")) - with patch.object(servicer._factory, "get_for_protocol", return_value=fake_connector): + with ( + patch.object(servicer._factory, "is_exposed", return_value=True), + patch.object(servicer._factory, "get", new=AsyncMock(return_value=fake_connector)), + ): req = connector_pb2.InvokeRequest(connector_id="x", action="act") resp = await servicer._invoke_async(req) assert resp.success is True @@ -329,10 +347,61 @@ async def test_invoke_error_response_maps_error_category( trace_id="t5", ) ) - with patch.object(servicer._factory, "get_for_protocol", return_value=fake_connector): + with ( + patch.object(servicer._factory, "is_exposed", return_value=True), + patch.object(servicer._factory, "get", new=AsyncMock(return_value=fake_connector)), + ): req = connector_pb2.InvokeRequest(connector_id="x", action="act", payload_json="{}") resp = await servicer._invoke_async(req) assert resp.success is False assert resp.error_code == "UPSTREAM_TIMEOUT" assert resp.error_category == ErrorCategory.RETRYABLE.value + + +async def test_invoke_missing_tenant_when_multitenancy_enabled( + servicer: ConnectorServiceServicer, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + with patch("bindings.grpc_server.server.get_grpc_caller_identity", return_value=None): + req = connector_pb2.InvokeRequest(connector_id="x", action="act") + resp = await servicer._invoke_async(req) + + assert resp.success is False + assert resp.error_code == "MISSING_TENANT" + assert resp.error_category == ErrorCategory.AUTH.value + + +async def test_invoke_config_not_found(servicer: ConnectorServiceServicer) -> None: + from node_wire_runtime.config_store import ConfigNotFoundError + + with ( + patch.object(servicer._factory, "is_exposed", return_value=True), + patch.object( + servicer._factory, + "get", + new=AsyncMock(side_effect=ConfigNotFoundError("no config")), + ), + ): + req = connector_pb2.InvokeRequest(connector_id="x", action="act", payload_json="{}") + resp = await servicer._invoke_async(req) + + assert resp.success is False + assert resp.error_code == "CONFIG_NOT_FOUND" + assert resp.error_category == ErrorCategory.AUTH.value + + +def test_invoke_passes_metadata_headers(servicer: ConnectorServiceServicer) -> None: + context = MagicMock() + context.invocation_metadata.return_value = (("x-tenant-id", "acme"),) + req = connector_pb2.InvokeRequest(connector_id="x", action="act") + + with patch("bindings.grpc_server.server._async_runner") as runner: + runner.run.return_value = connector_pb2.InvokeResponse(success=True, trace_id="t") + resp = servicer.Invoke(req, context) + + assert resp.success is True + call_args = runner.run.call_args[0][0] + # The coroutine was created with metadata; force close to avoid warnings. + call_args.close() diff --git a/tests/test_mcp_auth.py b/tests/test_mcp_auth.py index 97a7173e..bcf84478 100644 --- a/tests/test_mcp_auth.py +++ b/tests/test_mcp_auth.py @@ -129,6 +129,11 @@ async def test_mcp_authz_denies_tool_without_scope(monkeypatch: pytest.MonkeyPat assert identity is not None server = McpServer(connector_ids=["smtp"]) + # Rev4: configured = entitled. The JWT tenant must have a config to reach the + # scope policy (otherwise it fails closed at config resolution first). + server._factory.store.create( + "tenant-a", "smtp", {"name": "default", "default": True, "config": {}} + ) resp = await server.invoke_tool( "smtp.send_email", { @@ -149,10 +154,12 @@ async def test_mcp_authz_denies_tool_without_scope(monkeypatch: pytest.MonkeyPat async def test_mcp_execution_passes_principal_and_tenant( monkeypatch: pytest.MonkeyPatch, ) -> None: + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") monkeypatch.delenv("NW_MCP_AUTH_DISABLED", raising=False) monkeypatch.delenv("NW_MCP_API_KEY", raising=False) monkeypatch.setenv("NW_MCP_JWT_SECRET", "jwt-secret") monkeypatch.delenv("NW_MCP_ACTION_SCOPE_MAP_JSON", raising=False) + monkeypatch.delenv("NW_TENANT_ID", raising=False) token = mint_test_jwt( {"sub": "service-account", "tenant_id": "tenant-42", "scopes": ["*"]}, @@ -162,7 +169,12 @@ async def test_mcp_execution_passes_principal_and_tenant( assert identity is not None server = McpServer(connector_ids=["smtp"]) - smtp = server._factory.get_for_protocol("smtp", "mcp") + # Rev4: configured = entitled. Provision the JWT tenant's config and patch the + # tenant-scoped instance that invoke_tool will resolve. + server._factory.store.create( + "tenant-42", "smtp", {"name": "default", "default": True, "config": {}} + ) + smtp = await server._factory.get("smtp", tenant_id="tenant-42") assert smtp is not None captured: dict[str, object] = {} diff --git a/tests/test_mcp_multitenant.py b/tests/test_mcp_multitenant.py new file mode 100644 index 00000000..5dde32dc --- /dev/null +++ b/tests/test_mcp_multitenant.py @@ -0,0 +1,235 @@ +# +# SPDX-FileCopyrightText: 2026 AOT Technologies +# SPDX-License-Identifier: Apache-2.0 +# +"""MCP binding multitenancy: session pin, stdio env pin, config_name (§6.3 / §6.4).""" + +from __future__ import annotations + +import os + +import pytest + +from bindings.mcp_server.server import ( + McpServer, + _http_request_headers, + _session_tenant_ctx, +) +from node_wire_runtime.models import ConnectorResponse + + +@pytest.fixture(autouse=True) +def _mcp_mt_env(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv( + "NW_ALLOWED_CONNECTORS", + "http_generic,smtp,stripe,google_drive,fhir_epic,fhir_cerner", + ) + monkeypatch.setenv("NW_MCP_SCOPE_POLICY_DEFAULT", "allow") + monkeypatch.setenv("NW_MCP_AUTH_DISABLED", "true") + monkeypatch.delenv("NW_MCP_ACTION_SCOPE_MAP_JSON", raising=False) + monkeypatch.delenv("NW_MCP_API_KEY_SCOPES", raising=False) + monkeypatch.delenv("NW_TENANT_ID", raising=False) + monkeypatch.setenv("NW_RATE_LIMIT_DISABLED", "true") + + +async def _capture_invoke( + server: McpServer, + *, + tenant_for_instance: str, + arguments: dict | None = None, +) -> dict[str, object]: + """Provision tenant config, patch run(), invoke tool, return captured kwargs.""" + server._factory.store.create( + tenant_for_instance, + "http_generic", + {"name": "default", "default": True, "config": {}, "auth": {}}, + ) + connector = await server._factory.get( + "http_generic", tenant_id=tenant_for_instance, config_name="default" + ) + captured: dict[str, object] = {} + + async def fake_run(raw_input, *, principal=None, tenant_id=None, scopes=None): + captured["payload"] = dict(raw_input) + captured["tenant_id"] = tenant_id + return ConnectorResponse(success=True, data={"ok": True}, trace_id="t") + + orig = connector.run + connector.run = fake_run # type: ignore[method-assign] + try: + await server.invoke_tool( + "http_generic.request", + arguments + or { + "method": "GET", + "url": "https://example.com", + }, + ) + finally: + connector.run = orig # type: ignore[method-assign] + return captured + + +@pytest.mark.asyncio +async def test_mt_off_uses_default_tenant(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "false") + server = McpServer(connector_ids=["http_generic"]) + tools = server.list_tools() + assert all("config_name" not in (t["input_schema"].get("properties") or {}) for t in tools) + + connector = await server._factory.get("http_generic", tenant_id="__default__") + captured: dict[str, object] = {} + + async def fake_run(raw_input, *, principal=None, tenant_id=None, scopes=None): + captured["tenant_id"] = tenant_id + captured["payload"] = dict(raw_input) + return ConnectorResponse(success=True, data={"ok": True}, trace_id="t") + + orig = connector.run + connector.run = fake_run # type: ignore[method-assign] + hdr_tok = _http_request_headers.set({"x-tenant-id": "acme"}) + try: + await server.invoke_tool( + "http_generic.request", + { + "method": "GET", + "url": "https://example.com", + "config_name": "should-be-ignored", + }, + ) + finally: + connector.run = orig # type: ignore[method-assign] + _http_request_headers.reset(hdr_tok) + + assert captured["tenant_id"] == "__default__" + assert "config_name" not in captured["payload"] # type: ignore[operator] + + +@pytest.mark.asyncio +async def test_mt_on_stdio_env_pin_selects_tenant(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + server = McpServer(connector_ids=["http_generic"]) + server._stdio_env_tenant_pin = "acme" + + captured = await _capture_invoke(server, tenant_for_instance="acme") + assert captured["tenant_id"] == "acme" + assert "config_name" not in captured["payload"] # type: ignore[operator] + + +@pytest.mark.asyncio +async def test_mt_on_session_pin_ignores_process_nw_tenant_id( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + monkeypatch.setenv("NW_TENANT_ID", "from-env-must-not-win") + server = McpServer(connector_ids=["http_generic"]) + server._stdio_env_tenant_pin = "from-env-must-not-win" + + sess = _session_tenant_ctx.set("acme") + try: + captured = await _capture_invoke(server, tenant_for_instance="acme") + finally: + _session_tenant_ctx.reset(sess) + + assert captured["tenant_id"] == "acme" + + +@pytest.mark.asyncio +async def test_mt_on_missing_tenant_raises(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + server = McpServer(connector_ids=["http_generic"]) + server._stdio_env_tenant_pin = None + + with pytest.raises(ValueError, match="X-Tenant-ID is required"): + await server.invoke_tool( + "http_generic.request", + {"method": "GET", "url": "https://example.com"}, + ) + + +@pytest.mark.asyncio +async def test_mt_on_config_name_stripped_and_unknown_fail_closed( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + server = McpServer(connector_ids=["http_generic"]) + server._stdio_env_tenant_pin = "acme" + server._factory.store.create( + "acme", + "http_generic", + {"name": "primary", "default": True, "config": {}, "auth": {}}, + ) + + connector = await server._factory.get("http_generic", tenant_id="acme", config_name="primary") + captured: dict[str, object] = {} + + async def fake_run(raw_input, *, principal=None, tenant_id=None, scopes=None): + captured["payload"] = dict(raw_input) + return ConnectorResponse(success=True, data={"ok": True}, trace_id="t") + + connector.run = fake_run # type: ignore[method-assign] + await server.invoke_tool( + "http_generic.request", + { + "method": "GET", + "url": "https://example.com", + "config_name": "primary", + }, + ) + assert "config_name" not in captured["payload"] # type: ignore[operator] + + with pytest.raises(ValueError, match="not available via MCP"): + await server.invoke_tool( + "http_generic.request", + { + "method": "GET", + "url": "https://example.com", + "config_name": "does-not-exist", + }, + ) + + +@pytest.mark.asyncio +async def test_mt_on_config_name_in_tool_schema(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + server = McpServer(connector_ids=["http_generic"]) + tools = server.list_tools() + cfg_schemas = [ + (t["input_schema"].get("properties") or {}).get("config_name") + for t in tools + if "config_name" in (t["input_schema"].get("properties") or {}) + ] + assert cfg_schemas + assert cfg_schemas[0]["type"] == ["string", "null"] + + +@pytest.mark.asyncio +async def test_mt_on_null_config_name_uses_tenant_default( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """LLMs emit config_name: null; treat as omit (tenant default).""" + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + server = McpServer(connector_ids=["http_generic"]) + server._stdio_env_tenant_pin = "acme" + captured = await _capture_invoke( + server, + tenant_for_instance="acme", + arguments={ + "method": "GET", + "url": "https://example.com", + "config_name": None, + "headers": None, + }, + ) + assert captured["tenant_id"] == "acme" + assert "config_name" not in captured["payload"] # type: ignore[operator] + assert "headers" not in captured["payload"] # type: ignore[operator] + + +def test_stdio_pin_assignment_from_nw_tenant_id(monkeypatch: pytest.MonkeyPatch) -> None: + """Same assignment used at the start of _run_stdio_async.""" + monkeypatch.setenv("NW_TENANT_ID", " acme ") + server = McpServer(connector_ids=["http_generic"]) + raw = os.getenv("NW_TENANT_ID") + server._stdio_env_tenant_pin = raw.strip() if raw and raw.strip() else None + assert server._stdio_env_tenant_pin == "acme" diff --git a/tests/test_multitenancy.py b/tests/test_multitenancy.py new file mode 100644 index 00000000..7ba56141 --- /dev/null +++ b/tests/test_multitenancy.py @@ -0,0 +1,617 @@ +# +# SPDX-FileCopyrightText: 2026 AOT Technologies +# SPDX-License-Identifier: Apache-2.0 +# +"""Header-based tenancy and runtime config API (config store, factory, identity).""" + +from __future__ import annotations + +import threading +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi.testclient import TestClient +from pydantic import BaseModel + +from bindings.factory import ConnectorFactory +from bindings.rest_api.app import app, get_factory +from node_wire_runtime import BaseConnector, sdk_action +from node_wire_runtime.config_store import ( + ConfigNameConflictError, + ConfigNotFoundError, + ConnectorConfigStore, + DefaultDeletionError, + redact, +) +from node_wire_runtime.identity import ( + MissingTenantError, + is_multitenancy_enabled, + resolve_config_name, + resolve_tenant_id, + tenant_from_headers, +) +from node_wire_runtime.models import ConnectorResponse +from node_wire_runtime.secrets import ( + EnvSecretProvider, + TenantSecretNotFoundError, + TenantSecretProvider, +) + +# --------------------------------------------------------------------------- # +# Config store lifecycle +# --------------------------------------------------------------------------- # + + +def _doc(name: str, default: bool | None = None, **config): + d: dict = {"name": name, "config": config} + if default is not None: + d["default"] = default + return d + + +def test_init_first_is_default_when_none_marked(): + store = ConnectorConfigStore() + store.init({"acme": {"slack": [_doc("a"), _doc("b")]}}) + assert store.resolve("acme", "slack", None).name == "a" + + +def test_init_honours_explicit_default(): + store = ConnectorConfigStore() + store.init({"acme": {"slack": [_doc("a"), _doc("b", default=True)]}}) + assert store.resolve("acme", "slack", None).name == "b" + + +def test_init_rejects_two_defaults(): + store = ConnectorConfigStore() + with pytest.raises(Exception): + store.init({"acme": {"slack": [_doc("a", default=True), _doc("b", default=True)]}}) + + +def test_create_first_config_auto_defaults(): + store = ConnectorConfigStore() + rec = store.create("acme", "slack", _doc("only")) + assert rec.default is True + assert store.resolve("acme", "slack", None).name == "only" + + +def test_create_duplicate_name_conflicts(): + store = ConnectorConfigStore() + store.create("acme", "slack", _doc("a")) + with pytest.raises(ConfigNameConflictError): + store.create("acme", "slack", _doc("a")) + + +def test_delete_default_with_siblings_requires_new_default(): + store = ConnectorConfigStore() + store.create("acme", "slack", _doc("a")) + store.create("acme", "slack", _doc("b")) + with pytest.raises(DefaultDeletionError): + store.delete("acme", "slack", "a") # 'a' is default, 'b' remains + + +def test_delete_default_moves_flag_with_new_default(): + store = ConnectorConfigStore() + store.create("acme", "slack", _doc("a")) + store.create("acme", "slack", _doc("b")) + store.delete("acme", "slack", "a", new_default="b") + assert store.resolve("acme", "slack", None).name == "b" + + +def test_delete_last_config_removes_scope(): + store = ConnectorConfigStore() + store.create("acme", "slack", _doc("a")) + store.delete("acme", "slack", "a") # last config: allowed without new_default + assert store.has_config("acme", "slack") is False + with pytest.raises(ConfigNotFoundError): + store.resolve("acme", "slack", None) + + +def test_update_name_is_immutable(): + store = ConnectorConfigStore() + store.create("acme", "slack", _doc("a")) + with pytest.raises(Exception): + store.update("acme", "slack", "a", _doc("renamed")) + + +def test_set_default_moves_exactly_one_flag(): + store = ConnectorConfigStore() + store.create("acme", "slack", _doc("a")) + store.create("acme", "slack", _doc("b")) + store.set_default("acme", "slack", "b") + docs = {d["name"]: d["default"] for d in store.list("acme", "slack")} + assert docs == {"a": False, "b": True} + + +def test_redaction_masks_inline_values_on_read(): + store = ConnectorConfigStore() + store.create( + "acme", + "slack", + { + "name": "a", + "auth": {"provider": "static_token", "token_value": "xoxb-secret"}, + }, + ) + got = store.get("acme", "slack", "a") + assert got["auth"]["token_value"] != "xoxb-secret" + # The internal resolve path still sees the real value. + assert store.resolve("acme", "slack", "a").raw["auth"]["token_value"] == "xoxb-secret" + + +def test_redact_passes_references_through(): + out = redact({"auth": {"secret_key": "announcement_token", "password": "p"}}) + assert out["auth"]["secret_key"] == "announcement_token" + assert out["auth"]["password"] != "p" + + +# --------------------------------------------------------------------------- # +# Identity +# --------------------------------------------------------------------------- # + + +def test_tenant_from_headers_case_insensitive(): + assert tenant_from_headers({"X-Tenant-ID": "acme"}) == "acme" + assert tenant_from_headers({"x-tenant-id": "acme"}) == "acme" + + +def test_missing_header_resolves_default(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "false") + assert resolve_tenant_id(headers={}) == "__default__" + assert resolve_tenant_id(headers={"X-Tenant-ID": " "}) == "__default__" + + +def test_header_override_env(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr("node_wire_runtime.identity.TENANT_HEADER", "x-org-id") + assert tenant_from_headers({"X-Org-ID": "globex"}) == "globex" + + +def test_jwt_fallback_when_no_header(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + ident = MagicMock() + ident.tenant_id = "t-1" + assert resolve_tenant_id(headers={}, jwt_identity=ident) == "t-1" + + +def test_header_wins_over_jwt(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + ident = MagicMock() + ident.tenant_id = "t-1" + assert resolve_tenant_id(headers={"X-Tenant-ID": "acme"}, jwt_identity=ident) == "acme" + + +def test_env_pin_wins_over_everything(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + ident = MagicMock() + ident.tenant_id = "t-1" + assert ( + resolve_tenant_id( + headers={"X-Tenant-ID": "acme"}, jwt_identity=ident, env_pin="stdio-tenant" + ) + == "stdio-tenant" + ) + + +# --------------------------------------------------------------------------- # +# TenantSecretProvider +# --------------------------------------------------------------------------- # + + +def test_tenant_secret_provider_env_translation(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_ACME_SLACK_ANNOUNCEMENT_TOKEN", "xoxb-1") + provider = TenantSecretProvider(EnvSecretProvider(), "acme", "slack") + assert provider.get_secret("announcement_token") == "xoxb-1" + + +def test_tenant_secret_provider_with_config_name(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_ACME_SLACK_DEMO_ANNOUNCEMENT_TOKEN", "xoxb-cfg") + provider = TenantSecretProvider( + EnvSecretProvider(), "acme", "slack", config_name="demo" + ) + assert provider.get_secret("announcement_token") == "xoxb-cfg" + + +def test_tenant_secret_provider_is_strict(monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("NW_ACME_SLACK_MISSING", raising=False) + provider = TenantSecretProvider(EnvSecretProvider(), "acme", "slack") + with pytest.raises(TenantSecretNotFoundError): + provider.get_secret("missing") + + +# --------------------------------------------------------------------------- # +# Factory: three-part key, default routing, invalidation +# --------------------------------------------------------------------------- # + + +def _bare_factory(monkeypatch: pytest.MonkeyPatch) -> ConnectorFactory: + """Factory whose instantiation returns a fresh stub per call (no real connector).""" + factory = ConnectorFactory() + monkeypatch.setattr(factory, "_instantiate", lambda record: MagicMock(spec=BaseConnector)) + return factory + + +async def test_factory_three_part_key_isolates_instances(monkeypatch: pytest.MonkeyPatch): + factory = _bare_factory(monkeypatch) + factory.store.init( + { + "acme": {"slack": [_doc("internal", default=True), _doc("announce")]}, + "globex": {"slack": [_doc("main")]}, + } + ) + a1 = await factory.get("slack", tenant_id="acme", config_name="internal") + a2 = await factory.get("slack", tenant_id="acme", config_name="announce") + g1 = await factory.get("slack", tenant_id="globex", config_name="main") + assert a1 is not a2 # two configs of one connector -> distinct instances + assert a1 is not g1 # two tenants -> distinct instances + + +async def test_default_and_explicit_share_one_instance(monkeypatch: pytest.MonkeyPatch): + factory = _bare_factory(monkeypatch) + factory.store.init({"acme": {"slack": [_doc("internal", default=True), _doc("announce")]}}) + via_default = await factory.get("slack", tenant_id="acme", config_name=None) + via_name = await factory.get("slack", tenant_id="acme", config_name="internal") + assert via_default is via_name + + +async def test_moving_default_reroutes_without_duplicating(monkeypatch: pytest.MonkeyPatch): + factory = _bare_factory(monkeypatch) + factory.store.init({"acme": {"slack": [_doc("internal", default=True), _doc("announce")]}}) + first_default = await factory.get("slack", tenant_id="acme") + factory.store.set_default("acme", "slack", "announce") + new_default = await factory.get("slack", tenant_id="acme") + explicit_announce = await factory.get("slack", tenant_id="acme", config_name="announce") + assert new_default is not first_default + assert new_default is explicit_announce + + +async def test_write_invalidates_cached_instance(monkeypatch: pytest.MonkeyPatch): + factory = _bare_factory(monkeypatch) + factory.store.init({"acme": {"slack": [_doc("internal", default=True)]}}) + before = await factory.get("slack", tenant_id="acme", config_name="internal") + factory.store.update("acme", "slack", "internal", _doc("internal", channel="#new")) + after = await factory.get("slack", tenant_id="acme", config_name="internal") + assert before is not after # write evicted the cached instance + + +async def test_off_loop_write_does_not_raise(monkeypatch: pytest.MonkeyPatch): + factory = _bare_factory(monkeypatch) + factory.store.init({"acme": {"slack": [_doc("a", default=True)]}}) + await factory.get("slack", tenant_id="acme", config_name="a") + + errors: list[Exception] = [] + + def _mutate() -> None: + try: + factory.store.delete("acme", "slack", "a") + except Exception as exc: # noqa: BLE001 + errors.append(exc) + + t = threading.Thread(target=_mutate) + t.start() + t.join() + assert errors == [] + assert factory.store.has_config("acme", "slack") is False + + +# --------------------------------------------------------------------------- # +# Fail-closed +# --------------------------------------------------------------------------- # + + +async def test_unconfigured_and_unknown_name_both_raise(monkeypatch: pytest.MonkeyPatch): + factory = _bare_factory(monkeypatch) + factory.store.init({"acme": {"slack": [_doc("a", default=True)]}}) + with pytest.raises(ConfigNotFoundError): + await factory.get("slack", tenant_id="acme", config_name="does-not-exist") + with pytest.raises(ConfigNotFoundError): + await factory.get("slack", tenant_id="nobody") + + +def test_rest_fail_closed_returns_indistinguishable_403(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "false") + mock_factory = MagicMock() + mock_factory.is_exposed.return_value = True + mock_factory.get = AsyncMock(side_effect=ConfigNotFoundError("secret internals")) + app.dependency_overrides[get_factory] = lambda: mock_factory + try: + client = TestClient(app) + unknown_scope = client.post("/connectors/http_generic/request", json={}) + unknown_name = client.post("/connectors/http_generic/request", json={"config_name": "nope"}) + finally: + app.dependency_overrides.clear() + assert unknown_scope.status_code == 403 + assert unknown_name.status_code == 403 + # Same body: the internal reason is never leaked, so names cannot be enumerated. + assert unknown_scope.json() == unknown_name.json() + assert "secret internals" not in unknown_scope.text + + +def test_rest_missing_tenant_returns_400_when_multitenancy_enabled( + monkeypatch: pytest.MonkeyPatch, +): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + mock_factory = MagicMock() + mock_factory.is_exposed.return_value = True + mock_factory.get = AsyncMock() + app.dependency_overrides[get_factory] = lambda: mock_factory + try: + client = TestClient(app) + resp = client.post("/connectors/http_generic/request", json={}) + finally: + app.dependency_overrides.clear() + assert resp.status_code == 400 + assert "X-Tenant-ID is required" in resp.json()["detail"] + mock_factory.get.assert_not_called() + + +# --------------------------------------------------------------------------- # +# Direct integration (no bindings) — the library-framing acceptance test +# --------------------------------------------------------------------------- # + + +class _EchoIn(BaseModel): + action: str = "echo" + text: str = "" + + +class _EchoOut(BaseModel): + text: str = "" + channel: str = "" + + +class _EchoConnector(BaseConnector): + connector_id = "test_echo" + output_model = _EchoOut + + @sdk_action("echo", requires_auth=False) + async def echo(self, params: _EchoIn, *, trace_id: str) -> _EchoOut: + return _EchoOut(text=params.text, channel=self.config.get("channel", "")) + + +async def test_direct_integration_store_factory_run(monkeypatch: pytest.MonkeyPatch): + factory = ConnectorFactory() + factory.store.init( + { + "acme": { + "test_echo": [{"name": "primary", "default": True, "config": {"channel": "#eng"}}] + } + } + ) + connector = await factory.get("test_echo", tenant_id="acme") + assert connector._config_name == "primary" + assert connector.config == {"channel": "#eng"} + + resp: ConnectorResponse = await connector.run( + {"action": "echo", "text": "hi"}, tenant_id="acme" + ) + assert resp.success is True + assert resp.data["text"] == "hi" + assert resp.data["channel"] == "#eng" # per-config injection reached the connector + + +# --------------------------------------------------------------------------- # +# NW_MULTITENANCY_ENABLED feature flag +# --------------------------------------------------------------------------- # + + +def test_multitenancy_disabled_by_default(monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("NW_MULTITENANCY_ENABLED", raising=False) + assert is_multitenancy_enabled() is False + + +def test_multitenancy_enabled_truthy_values(monkeypatch: pytest.MonkeyPatch): + for val in ("true", "1", "yes", "on", "True", "YES"): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", val) + assert is_multitenancy_enabled() is True, f"Expected True for {val!r}" + + +def test_multitenancy_disabled_falsy_values(monkeypatch: pytest.MonkeyPatch): + for val in ("false", "0", "no", "off", "False"): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", val) + assert is_multitenancy_enabled() is False, f"Expected False for {val!r}" + + +def test_resolve_tenant_id_disabled_ignores_header(monkeypatch: pytest.MonkeyPatch): + """When disabled, resolve_tenant_id always returns DEFAULT_TENANT.""" + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "false") + assert resolve_tenant_id(headers={"X-Tenant-ID": "acme"}) == "__default__" + + +def test_resolve_tenant_id_disabled_ignores_jwt(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "false") + ident = MagicMock() + ident.tenant_id = "t-1" + assert resolve_tenant_id(headers={}, jwt_identity=ident) == "__default__" + + +def test_resolve_tenant_id_disabled_ignores_env_pin(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "false") + assert resolve_tenant_id(env_pin="stdio-tenant") == "__default__" + + +def test_resolve_tenant_id_enabled_reads_header(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + assert resolve_tenant_id(headers={"X-Tenant-ID": "acme"}) == "acme" + + +def test_resolve_tenant_id_enabled_allows_explicit_default(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + assert resolve_tenant_id(headers={"X-Tenant-ID": "__default__"}) == "__default__" + + +def test_resolve_tenant_id_enabled_missing_raises(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + with pytest.raises(MissingTenantError, match="X-Tenant-ID is required"): + resolve_tenant_id(headers={}) + with pytest.raises(MissingTenantError): + resolve_tenant_id(headers={"X-Tenant-ID": " "}) + + +def test_resolve_tenant_id_enabled_env_pin_wins(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + assert ( + resolve_tenant_id(headers={"X-Tenant-ID": "acme"}, env_pin="stdio-tenant") == "stdio-tenant" + ) + + +def test_resolve_config_name_disabled_returns_none(monkeypatch: pytest.MonkeyPatch): + """When multitenancy is off, user-supplied config names are suppressed.""" + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "false") + assert resolve_config_name("my-config") is None + assert resolve_config_name(None) is None + + +def test_resolve_config_name_enabled_passthrough(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + assert resolve_config_name("my-config") == "my-config" + assert resolve_config_name(None) is None + assert resolve_config_name("") is None + assert resolve_config_name(" ") is None + + +def test_resolve_config_name_enabled_rejects_non_string(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + # JSON null and LLM quirks must map to omit (tenant default), not fail closed. + assert resolve_config_name(None) is None # type: ignore[arg-type] + + +# --------------------------------------------------------------------------- # +# REST config CRUD (patch coverage for bindings.rest_api.app) +# --------------------------------------------------------------------------- # + + +def _config_factory() -> ConnectorFactory: + factory = ConnectorFactory() + # No YAML load needed — store starts empty; CRUD tests seed via API/store. + return factory + + +def test_rest_config_crud_roundtrip(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + factory = _config_factory() + app.dependency_overrides[get_factory] = lambda: factory + headers = {"X-Tenant-ID": "acme"} + try: + client = TestClient(app) + + create = client.post( + "/v1/connectors/slack/configs", + json={"name": "primary", "default": True, "config": {"channel": "#a"}}, + headers=headers, + ) + assert create.status_code == 201 + assert create.json() == {"name": "primary", "default": True} + + listed = client.get("/v1/connectors/slack/configs", headers=headers) + assert listed.status_code == 200 + assert any(item["name"] == "primary" for item in listed.json()) + + got = client.get("/v1/connectors/slack/configs/primary", headers=headers) + assert got.status_code == 200 + assert got.json()["name"] == "primary" + + client.post( + "/v1/connectors/slack/configs", + json={"name": "secondary", "config": {"channel": "#b"}}, + headers=headers, + ) + updated = client.put( + "/v1/connectors/slack/configs/secondary", + json={"name": "secondary", "config": {"channel": "#b2"}}, + headers=headers, + ) + assert updated.status_code == 200 + + set_def = client.put( + "/v1/connectors/slack/configs/secondary/default", + headers=headers, + ) + assert set_def.status_code == 200 + + deleted = client.delete( + "/v1/connectors/slack/configs/primary", + headers=headers, + ) + assert deleted.status_code == 200 + + missing = client.get("/v1/connectors/slack/configs/primary", headers=headers) + assert missing.status_code == 404 + finally: + app.dependency_overrides.clear() + + +def test_rest_config_missing_tenant_returns_400(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + factory = _config_factory() + app.dependency_overrides[get_factory] = lambda: factory + try: + client = TestClient(app) + resp = client.post( + "/v1/connectors/slack/configs", + json={"name": "x", "config": {}}, + ) + finally: + app.dependency_overrides.clear() + assert resp.status_code == 400 + assert "X-Tenant-ID" in resp.json()["detail"] + + +def test_rest_config_create_upserts_on_conflict(monkeypatch: pytest.MonkeyPatch): + """POST create is create-or-update for playground per-connector Add config.""" + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + factory = _config_factory() + factory.store.create("acme", "slack", _doc("dup", default=True)) + app.dependency_overrides[get_factory] = lambda: factory + try: + client = TestClient(app) + resp = client.post( + "/v1/connectors/slack/configs", + json={"name": "dup", "config": {"channel": "#updated"}}, + headers={"X-Tenant-ID": "acme"}, + ) + finally: + app.dependency_overrides.clear() + assert resp.status_code == 201 + assert resp.json()["name"] == "dup" + got = factory.store.get("acme", "slack", "dup") + assert got is not None + assert got.get("config", {}).get("channel") == "#updated" + + +def test_rest_config_init_with_tenant_header(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + factory = _config_factory() + app.dependency_overrides[get_factory] = lambda: factory + try: + client = TestClient(app) + resp = client.post( + "/v1/config/init", + json={"slack": [{"name": "boot", "default": True, "config": {}}]}, + headers={"X-Tenant-ID": "acme"}, + ) + finally: + app.dependency_overrides.clear() + assert resp.status_code == 200 + assert factory.store.resolve("acme", "slack", None).name == "boot" + + +def test_rest_config_delete_default_without_replacement_returns_400( + monkeypatch: pytest.MonkeyPatch, +): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + factory = _config_factory() + factory.store.create("acme", "slack", _doc("keep", default=True)) + factory.store.create("acme", "slack", _doc("other")) + app.dependency_overrides[get_factory] = lambda: factory + try: + client = TestClient(app) + resp = client.delete( + "/v1/connectors/slack/configs/keep", + headers={"X-Tenant-ID": "acme"}, + ) + finally: + app.dependency_overrides.clear() + assert resp.status_code == 400 + + +def test_tenant_from_headers_none_value(): + assert tenant_from_headers({"X-Tenant-ID": None}) is None # type: ignore[dict-item] diff --git a/tests/test_playground_agent_tenant.py b/tests/test_playground_agent_tenant.py new file mode 100644 index 00000000..6a78ed38 --- /dev/null +++ b/tests/test_playground_agent_tenant.py @@ -0,0 +1,113 @@ +# +# SPDX-FileCopyrightText: 2026 AOT Technologies +# SPDX-License-Identifier: Apache-2.0 +# +"""Playground agent chat forwards tenant to MCP clients.""" + +from __future__ import annotations + +import json +from unittest.mock import MagicMock, patch + +import httpx +import pytest +from fastapi.testclient import TestClient + +from agents.toolhive import InProcessMcpClient, ToolHiveMcpClient + + +@pytest.fixture(autouse=True) +def _agent_tenant_env(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv( + "NW_ALLOWED_CONNECTORS", + "http_generic,smtp,stripe,google_drive,fhir_epic,fhir_cerner", + ) + monkeypatch.setenv("NW_REST_AUTH_DISABLED", "true") + monkeypatch.setenv("NW_MCP_SCOPE_POLICY_DEFAULT", "allow") + monkeypatch.setenv("NW_RATE_LIMIT_DISABLED", "true") + monkeypatch.delenv("TOOLHIVE_MCP_URL", raising=False) + monkeypatch.delenv("TOOLHIVE_MCP_URLS", raising=False) + monkeypatch.setenv("NW_MCP_TRANSPORT", "stdio") + + +@pytest.mark.asyncio +async def test_toolhive_client_sends_tenant_header() -> None: + seen: list[dict] = [] + + def handler(request: httpx.Request) -> httpx.Response: + seen.append({k.lower(): v for k, v in request.headers.items()}) + body = json.loads(request.content.decode()) + method = body.get("method") + req_id = body.get("id") + if method == "initialize": + return httpx.Response( + 200, + json={"jsonrpc": "2.0", "id": req_id, "result": {"protocolVersion": "2024-11-05"}}, + headers={"Mcp-Session-Id": "s1"}, + ) + if method == "notifications/initialized": + return httpx.Response(200, json={}) + if method == "tools/list": + return httpx.Response( + 200, + json={"jsonrpc": "2.0", "id": req_id, "result": {"tools": []}}, + ) + return httpx.Response(404) + + transport = httpx.MockTransport(handler) + _Real = httpx.AsyncClient + + def make_client(**kwargs: object) -> httpx.AsyncClient: + return _Real(transport=transport, timeout=60.0) + + with patch("httpx.AsyncClient", side_effect=make_client): + client = ToolHiveMcpClient( + "http://127.0.0.1:9/mcp", + extra_headers={"x-tenant-id": "acme"}, + ) + await client.list_tools() + + assert any(h.get("x-tenant-id") == "acme" for h in seen) + + +@pytest.mark.asyncio +async def test_inprocess_mcp_client_pins_tenant_and_config_name() -> None: + server = MagicMock() + server.list_tools.return_value = [{"name": "http_generic.request", "input_schema": {}}] + + async def fake_invoke(name, arguments, identity=None): + return {"ok": True, "args": arguments} + + server.invoke_tool = fake_invoke + + client = InProcessMcpClient(server, tenant_id="acme", config_name="primary") + async with client: + assert server._stdio_env_tenant_pin == "acme" + tools = await client.list_tools() + assert tools[0]["name"] == "http_generic.request" + out = await client.call_tool("http_generic.request", {"url": "https://example.com"}) + data = json.loads(out) + assert data["args"]["config_name"] == "primary" + assert data["args"]["url"] == "https://example.com" + + # Null from LLM must not block host-injected config_name. + out2 = await client.call_tool( + "http_generic.request", + {"url": "https://example.com", "config_name": None, "query": None}, + ) + data2 = json.loads(out2) + assert data2["args"]["config_name"] == "primary" + assert "query" not in data2["args"] + + +def test_agent_chat_requires_tenant_when_mt_on(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + from bindings.rest_api.app import app + + client = TestClient(app) + resp = client.post( + "/scenarios/agent-chat", + json={"message": "hello", "history": []}, + ) + assert resp.status_code == 400 + assert "X-Tenant-ID" in resp.json().get("detail", "") diff --git a/tests/test_rest_body_limit.py b/tests/test_rest_body_limit.py index 2590416b..9c0c473a 100644 --- a/tests/test_rest_body_limit.py +++ b/tests/test_rest_body_limit.py @@ -53,7 +53,8 @@ def test_rejects_oversized_content_length_on_connector_route() -> None: def test_allows_request_under_limit() -> None: mock_factory = MagicMock() - mock_factory.get_for_protocol.return_value = _stub_connector() + mock_factory.is_exposed.return_value = True + mock_factory.get = AsyncMock(return_value=_stub_connector()) app.dependency_overrides[get_factory] = lambda: mock_factory try: client = TestClient(app) diff --git a/tests/test_rest_connector_dispatch.py b/tests/test_rest_connector_dispatch.py index 31ceb97b..38c49888 100644 --- a/tests/test_rest_connector_dispatch.py +++ b/tests/test_rest_connector_dispatch.py @@ -107,8 +107,11 @@ def test_rest_dispatch_smoke_per_connector( payload: dict, ) -> None: mock_factory = MagicMock() - mock_factory.get_for_protocol.return_value = _stub_connector( - ConnectorResponse(success=True, data={"ok": True}, trace_id=f"t-{connector_id}") + mock_factory.is_exposed.return_value = True + mock_factory.get = AsyncMock( + return_value=_stub_connector( + ConnectorResponse(success=True, data={"ok": True}, trace_id=f"t-{connector_id}") + ) ) app.dependency_overrides[get_factory] = lambda: mock_factory try: @@ -119,7 +122,9 @@ def test_rest_dispatch_smoke_per_connector( assert response.status_code == 200 assert response.json()["success"] is True - mock_factory.get_for_protocol.assert_called_with(connector_id, "rest", action=action) + mock_factory.get.assert_awaited_with( + connector_id, tenant_id="__default__", config_name=None, action=action + ) @pytest.mark.parametrize( diff --git a/tests/test_rest_rate_limit_enforcement.py b/tests/test_rest_rate_limit_enforcement.py index c0ca6e3b..956a0aed 100644 --- a/tests/test_rest_rate_limit_enforcement.py +++ b/tests/test_rest_rate_limit_enforcement.py @@ -30,7 +30,8 @@ def _make_client(monkeypatch) -> tuple[TestClient, MagicMock]: monkeypatch.setattr(rest_app_module, "_rate_limiter_cfg", None) mock_factory = MagicMock() - mock_factory.get_for_protocol.return_value = _stub_connector() + mock_factory.is_exposed.return_value = True + mock_factory.get = AsyncMock(return_value=_stub_connector()) app.dependency_overrides[get_factory] = lambda: mock_factory return TestClient(app), mock_factory @@ -90,7 +91,8 @@ def test_rest_rate_limit_isolated_by_identity(monkeypatch) -> None: monkeypatch.setattr(rest_app_module, "_rate_limiter_cfg", None) mock_factory = MagicMock() - mock_factory.get_for_protocol.return_value = _stub_connector() + mock_factory.is_exposed.return_value = True + mock_factory.get = AsyncMock(return_value=_stub_connector()) app.dependency_overrides[get_factory] = lambda: mock_factory try: @@ -157,7 +159,8 @@ def test_rest_rate_limit_ignores_spoofed_xff_when_proxy_hops_zero(monkeypatch) - monkeypatch.setattr(rest_app_module, "_rate_limiter_cfg", None) mock_factory = MagicMock() - mock_factory.get_for_protocol.return_value = _stub_connector() + mock_factory.is_exposed.return_value = True + mock_factory.get = AsyncMock(return_value=_stub_connector()) app.dependency_overrides[get_factory] = lambda: mock_factory try: diff --git a/tests/test_tenant_store.py b/tests/test_tenant_store.py new file mode 100644 index 00000000..8b531f71 --- /dev/null +++ b/tests/test_tenant_store.py @@ -0,0 +1,405 @@ +# +# SPDX-FileCopyrightText: 2026 AOT Technologies +# SPDX-License-Identifier: Apache-2.0 +# +"""Tenant store persistence + per-config secrets overlay.""" + +from __future__ import annotations + +from pathlib import Path + +import pytest +import yaml +from fastapi.testclient import TestClient + +from bindings.factory import ConnectorFactory +from bindings.rest_api.app import app, get_factory +from bindings.rest_api.tenant_store import ( + load_tenants, + save_tenants, + upsert_tenant_secrets, +) +from node_wire_runtime.secrets import OverlaySecretProvider, tenant_scoped_secret_key + + +def _factory() -> ConnectorFactory: + return ConnectorFactory() + + +def test_overlay_resolves_before_env(monkeypatch: pytest.MonkeyPatch): + overlay = OverlaySecretProvider.instance() + key = tenant_scoped_secret_key( + "acme", "google_drive", "GOOGLE_DRIVE_SA_JSON", config_name="test" + ) + monkeypatch.setenv(key, "from-env") + overlay.set_secret(key, "from-overlay") + assert overlay.get_secret(key) == "from-overlay" + + +def test_drive_then_epic_same_tenant_config_name(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + monkeypatch.setenv("EPIC_FHIR_BASE_URL", "https://fhir.example/r4") + monkeypatch.setenv("EPIC_TOKEN_URL", "https://auth.example/token") + factory = _factory() + app.dependency_overrides[get_factory] = lambda: factory + headers = {"X-Tenant-ID": "acme"} + try: + client = TestClient(app) + + drive = client.post( + "/v1/connectors/google_drive/configs", + json={ + "name": "test", + "default": True, + "config": {}, + "auth": { + "provider": "service_account", + "sa_json_secret": "GOOGLE_DRIVE_SA_JSON", + "scopes": ["https://www.googleapis.com/auth/drive"], + }, + "secrets": {"GOOGLE_DRIVE_SA_JSON": '{"type":"service_account"}'}, + }, + headers=headers, + ) + assert drive.status_code == 201, drive.text + + epic = client.post( + "/v1/connectors/fhir_epic/configs", + json={ + "name": "test", + "default": True, + "config": {}, + "auth": { + "provider": "oauth2", + "grant_method": "private_key_jwt", + "token_url_secret": "EPIC_TOKEN_URL", + "client_id_secret": "EPIC_CLIENT_ID", + "private_key_secret": "EPIC_PRIVATE_KEY", + "kid_secret": "EPIC_KID", + "algorithm": "RS384", + }, + "secrets": { + "EPIC_CLIENT_ID": "cid", + "EPIC_PRIVATE_KEY": "-----BEGIN PRIVATE KEY-----\nabc\n-----END PRIVATE KEY-----", + "EPIC_KID": "kid-1", + }, + }, + headers=headers, + ) + assert epic.status_code == 201, epic.text + + drive_list = client.get("/v1/connectors/google_drive/configs", headers=headers) + epic_list = client.get("/v1/connectors/fhir_epic/configs", headers=headers) + assert any(c["name"] == "test" for c in drive_list.json()) + assert any(c["name"] == "test" for c in epic_list.json()) + + assert factory.store.has_config("acme", "google_drive") + assert factory.store.has_config("acme", "fhir_epic") + + tenants = client.get("/v1/tenants") + assert tenants.status_code == 200 + assert "acme" in tenants.json()["tenants"] + + keys = client.get( + "/v1/connectors/fhir_epic/secrets?config_name=test", + headers=headers, + ) + assert keys.status_code == 200 + key_list = keys.json()["keys"] + assert "EPIC_CLIENT_ID" in key_list + assert "EPIC_TOKEN_URL" in key_list + assert "epic_fhir_base_url" in key_list + assert "cid" not in str(keys.json()) + + assert not factory.store.has_config("acme", "slack") + finally: + app.dependency_overrides.clear() + + +def test_tenants_persist_roundtrip(monkeypatch: pytest.MonkeyPatch, tmp_path: Path): + path = tmp_path / "roundtrip.yaml" + monkeypatch.setenv("NW_TENANTS_PATH", str(path)) + monkeypatch.setenv("EPIC_TOKEN_URL", "https://token.example") + + store_factory = _factory() + upsert_tenant_secrets( + "acme", + "google_drive", + {"GOOGLE_DRIVE_SA_JSON": '{"type":"service_account","client_email":"a@b.c"}'}, + config_name="test", + ) + store_factory.store.create( + "acme", + "google_drive", + { + "name": "test", + "default": True, + "config": {}, + "auth": {"provider": "service_account", "sa_json_secret": "GOOGLE_DRIVE_SA_JSON"}, + }, + ) + save_tenants(store_factory.store) + assert path.is_file() + raw = yaml.safe_load(path.read_text(encoding="utf-8")) + assert "acme" in raw["tenants"] + assert raw["secrets"]["acme"]["google_drive"]["test"]["GOOGLE_DRIVE_SA_JSON"] == ( + '{"type":"service_account","client_email":"a@b.c"}' + ) + + OverlaySecretProvider.instance().clear() + fresh = _factory() + load_tenants(fresh.store) + assert fresh.store.has_config("acme", "google_drive") + scoped = tenant_scoped_secret_key( + "acme", "google_drive", "GOOGLE_DRIVE_SA_JSON", config_name="test" + ) + assert OverlaySecretProvider.instance().get_secret(scoped) == ( + '{"type":"service_account","client_email":"a@b.c"}' + ) + + +def test_new_config_without_secrets_rejected(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + factory = _factory() + app.dependency_overrides[get_factory] = lambda: factory + headers = {"X-Tenant-ID": "acme"} + try: + client = TestClient(app) + # Seed one config with secrets + first = client.post( + "/v1/connectors/google_drive/configs", + json={ + "name": "test", + "default": True, + "config": {}, + "auth": { + "provider": "service_account", + "sa_json_secret": "GOOGLE_DRIVE_SA_JSON", + }, + "secrets": {"GOOGLE_DRIVE_SA_JSON": '{"type":"service_account"}'}, + }, + headers=headers, + ) + assert first.status_code == 201, first.text + + # Sibling config cannot reuse those secrets when none are supplied + second = client.post( + "/v1/connectors/google_drive/configs", + json={ + "name": "test new 1", + "default": False, + "config": {}, + "auth": { + "provider": "service_account", + "sa_json_secret": "GOOGLE_DRIVE_SA_JSON", + }, + "secrets": {}, + }, + headers=headers, + ) + assert second.status_code == 400 + assert "GOOGLE_DRIVE_SA_JSON" in second.json()["detail"] + finally: + app.dependency_overrides.clear() + + +def test_per_config_secrets_are_isolated(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + factory = _factory() + app.dependency_overrides[get_factory] = lambda: factory + headers = {"X-Tenant-ID": "acme"} + sa_a = '{"type":"service_account","client_email":"a@x"}' + sa_b = '{"type":"service_account","client_email":"b@x"}' + try: + client = TestClient(app) + assert ( + client.post( + "/v1/connectors/google_drive/configs", + json={ + "name": "test", + "default": True, + "config": {}, + "auth": { + "provider": "service_account", + "sa_json_secret": "GOOGLE_DRIVE_SA_JSON", + }, + "secrets": {"GOOGLE_DRIVE_SA_JSON": sa_a}, + }, + headers=headers, + ).status_code + == 201 + ) + assert ( + client.post( + "/v1/connectors/google_drive/configs", + json={ + "name": "test new 1", + "default": False, + "config": {}, + "auth": { + "provider": "service_account", + "sa_json_secret": "GOOGLE_DRIVE_SA_JSON", + }, + "secrets": {"GOOGLE_DRIVE_SA_JSON": sa_b}, + }, + headers=headers, + ).status_code + == 201 + ) + + key_a = tenant_scoped_secret_key( + "acme", "google_drive", "GOOGLE_DRIVE_SA_JSON", config_name="test" + ) + key_b = tenant_scoped_secret_key( + "acme", "google_drive", "GOOGLE_DRIVE_SA_JSON", config_name="test new 1" + ) + overlay = OverlaySecretProvider.instance() + assert overlay.get_secret(key_a) == sa_a + assert overlay.get_secret(key_b) == sa_b + assert key_a != key_b + finally: + app.dependency_overrides.clear() + + +def test_secrets_put_rejects_missing_required(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + factory = _factory() + app.dependency_overrides[get_factory] = lambda: factory + try: + client = TestClient(app) + resp = client.put( + "/v1/connectors/fhir_epic/secrets", + json={"config_name": "demo", "secrets": {"EPIC_CLIENT_ID": "only-id"}}, + headers={"X-Tenant-ID": "acme"}, + ) + finally: + app.dependency_overrides.clear() + assert resp.status_code == 400 + assert "EPIC_PRIVATE_KEY" in resp.json()["detail"] + + +def test_secrets_put_rejects_bad_format(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + factory = _factory() + app.dependency_overrides[get_factory] = lambda: factory + try: + client = TestClient(app) + resp = client.put( + "/v1/connectors/slack/secrets", + json={"config_name": "demo", "secrets": {"SLACK_BOT_TOKEN": "not-a-bot-token"}}, + headers={"X-Tenant-ID": "acme"}, + ) + assert resp.status_code == 400 + assert "xoxb-" in resp.json()["detail"] + + pem_resp = client.put( + "/v1/connectors/fhir_epic/secrets", + json={ + "config_name": "demo", + "secrets": { + "EPIC_CLIENT_ID": "cid-1", + "EPIC_PRIVATE_KEY": "not-a-pem", + "EPIC_KID": "kid-1", + }, + }, + headers={"X-Tenant-ID": "acme"}, + ) + assert pem_resp.status_code == 400 + assert "PEM" in pem_resp.json()["detail"] + finally: + app.dependency_overrides.clear() + + +def test_secrets_partial_update_keeps_existing(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + monkeypatch.setenv("EPIC_TOKEN_URL", "https://token.example") + monkeypatch.setenv("EPIC_FHIR_BASE_URL", "https://fhir.example") + factory = _factory() + app.dependency_overrides[get_factory] = lambda: factory + headers = {"X-Tenant-ID": "acme"} + pem = "-----BEGIN PRIVATE KEY-----\nabc\n-----END PRIVATE KEY-----" + try: + client = TestClient(app) + first = client.put( + "/v1/connectors/fhir_epic/secrets", + json={ + "config_name": "demo", + "secrets": { + "EPIC_CLIENT_ID": "cid-1", + "EPIC_PRIVATE_KEY": pem, + "EPIC_KID": "kid-1", + }, + }, + headers=headers, + ) + assert first.status_code == 200 + partial = client.put( + "/v1/connectors/fhir_epic/secrets", + json={"config_name": "demo", "secrets": {"EPIC_KID": "kid-2"}}, + headers=headers, + ) + assert partial.status_code == 200 + keys = client.get( + "/v1/connectors/fhir_epic/secrets?config_name=demo", + headers=headers, + ) + assert "EPIC_CLIENT_ID" in keys.json()["keys"] + assert "EPIC_KID" in keys.json()["keys"] + scoped = tenant_scoped_secret_key( + "acme", "fhir_epic", "EPIC_CLIENT_ID", config_name="demo" + ) + assert OverlaySecretProvider.instance().get_secret(scoped) == "cid-1" + scoped_kid = tenant_scoped_secret_key( + "acme", "fhir_epic", "EPIC_KID", config_name="demo" + ) + assert OverlaySecretProvider.instance().get_secret(scoped_kid) == "kid-2" + finally: + app.dependency_overrides.clear() + + +def test_delete_config_clears_only_that_config_secrets(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NW_MULTITENANCY_ENABLED", "true") + factory = _factory() + app.dependency_overrides[get_factory] = lambda: factory + headers = {"X-Tenant-ID": "acme"} + try: + client = TestClient(app) + for name, token in (("keep", "xoxb-keep"), ("drop", "xoxb-drop")): + created = client.post( + "/v1/connectors/slack/configs", + json={ + "name": name, + "default": name == "keep", + "config": {}, + "auth": {"provider": "static_token", "secret_key": "SLACK_BOT_TOKEN"}, + "secrets": {"SLACK_BOT_TOKEN": token}, + }, + headers=headers, + ) + assert created.status_code == 201, created.text + + deleted = client.delete("/v1/connectors/slack/configs/drop", headers=headers) + assert deleted.status_code == 200 + + keep_keys = client.get( + "/v1/connectors/slack/secrets?config_name=keep", + headers=headers, + ) + assert "SLACK_BOT_TOKEN" in keep_keys.json()["keys"] + drop_keys = client.get( + "/v1/connectors/slack/secrets?config_name=drop", + headers=headers, + ) + assert drop_keys.json()["keys"] == [] + + keep_scoped = tenant_scoped_secret_key( + "acme", "slack", "SLACK_BOT_TOKEN", config_name="keep" + ) + drop_scoped = tenant_scoped_secret_key( + "acme", "slack", "SLACK_BOT_TOKEN", config_name="drop" + ) + assert OverlaySecretProvider.instance().get_secret(keep_scoped) == "xoxb-keep" + with pytest.raises(Exception): + OverlaySecretProvider.instance().get_secret(drop_scoped) + finally: + app.dependency_overrides.clear() diff --git a/tests/test_toolhive_agent.py b/tests/test_toolhive_agent.py index 7321146c..3727f356 100644 --- a/tests/test_toolhive_agent.py +++ b/tests/test_toolhive_agent.py @@ -26,15 +26,46 @@ LLMResponse, ToolCall, ) +from agents.schema_utils import openai_compatible_tool_parameters from agents.toolhive import ( ToolHiveAgent, ToolHiveMcpClient, _is_tool_failure, + omit_null_tool_args, resolve_max_tool_failures, truncate_tool_result_for_llm, ) +def test_omit_null_tool_args() -> None: + assert omit_null_tool_args({"a": 1, "b": None, "c": ""}) == {"a": 1, "c": ""} + + +def test_openai_compatible_tool_parameters_allows_null_on_optional() -> None: + schema = { + "type": "object", + "properties": { + "required_str": {"type": "string"}, + "optional_str": {"type": "string"}, + "already": {"type": ["string", "null"]}, + "list_no_null": {"type": ["string", "number"]}, + "null_only": {"type": "null"}, + "no_type": {"description": "untyped"}, + }, + "required": ["required_str"], + } + out = openai_compatible_tool_parameters(schema) + assert out["properties"]["required_str"]["type"] == "string" + assert out["properties"]["optional_str"]["type"] == ["string", "null"] + assert out["properties"]["already"]["type"] == ["string", "null"] + assert out["properties"]["list_no_null"]["type"] == ["string", "number", "null"] + assert out["properties"]["null_only"]["type"] == "null" + assert "type" not in out["properties"]["no_type"] + # Input schema must not be mutated. + assert schema["properties"]["optional_str"]["type"] == "string" + assert schema["properties"]["list_no_null"]["type"] == ["string", "number"] + + def test_truncate_tool_result_for_llm_respects_limit(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv("TOOLHIVE_MAX_TOOL_RESULT_CHARS", "20") long = "x" * 100