import { createCipheriv, createDecipheriv, createHash, createPublicKey, createVerify, constants, randomBytes, scryptSync } from "node:crypto"; import { createPasswordRecord, dbAll, dbGet, dbRun } from "./db.mjs"; import { createMfaChallenge, createMfaEnrollmentChallenge, createSession, mfaRequiredForUser, mfaStatus, safeUser } from "./auth.mjs"; const LOGIN_STATE_TTL_MS = 10 * 60 * 1000; const TICKET_TTL_MS = 60 * 1000; const CLOCK_SKEW_MS = 60 * 1000; const OIDC_STORAGE_KEY = scryptSync(process.env.AI_DRAMA_OIDC_STORAGE_KEY || process.env.AI_DRAMA_SESSION_SECRET || "ai-drama-local-oidc-key-v1", "ai-drama-oidc-storage", 32); const OIDC_JWKS_CACHE = new Map(); const { RSA_PKCS1_PSS_PADDING, RSA_PSS_SALTLEN_DIGEST } = constants; function oidcError(status, code, message, details = {}) { const error = new Error(message); error.status = status; error.code = code; error.details = details; return error; } function nowIso() { return new Date().toISOString(); } function hashValue(value) { return createHash("sha256").update(String(value || "")).digest("hex"); } function base64url(buffer) { return Buffer.from(buffer).toString("base64url"); } function decodeBase64url(value) { return Buffer.from(String(value || ""), "base64url"); } function encryptOpaque(value) { const iv = randomBytes(12); const cipher = createCipheriv("aes-256-gcm", OIDC_STORAGE_KEY, iv); const encrypted = Buffer.concat([cipher.update(String(value), "utf8"), cipher.final()]); return [iv, cipher.getAuthTag(), encrypted].map((part) => part.toString("base64url")).join("."); } function decryptOpaque(payload) { const [ivValue, tagValue, encryptedValue] = String(payload || "").split("."); if (!ivValue || !tagValue || !encryptedValue) throw oidcError(500, "oidc_state_invalid", "OIDC 登录状态存储记录无效"); const decipher = createDecipheriv("aes-256-gcm", OIDC_STORAGE_KEY, decodeBase64url(ivValue)); decipher.setAuthTag(decodeBase64url(tagValue)); return Buffer.concat([decipher.update(decodeBase64url(encryptedValue)), decipher.final()]).toString("utf8"); } function parseJson(value, fallback) { try { return JSON.parse(value); } catch { return fallback; } } function normalizeIssuer(value) { try { return new URL(String(value || "")).toString().replace(/\/+$/, ""); } catch { return String(value || "").replace(/\/+$/, ""); } } function claimValue(claims, path) { const keys = String(path || "").split(".").filter(Boolean); let current = claims; for (const key of keys) { if (!current || typeof current !== "object") return ""; current = current[key]; } return Array.isArray(current) ? current[0] : current; } function providerRow(providerId) { const provider = dbGet("SELECT * FROM identity_providers WHERE id = ? AND enabled = 1", [providerId]); if (!provider) throw oidcError(404, "sso_provider_not_found", "企业身份提供商不存在或未启用"); if (provider.kind !== "oidc") throw oidcError(400, "sso_provider_kind_unsupported", "当前登录链路只支持 OIDC,SAML 仍需部署层断言消费适配"); if (!provider.issuer_url || !provider.client_id) throw oidcError(400, "sso_provider_incomplete", "OIDC 提供商缺少 Issuer 或 Client ID"); return provider; } function identityPolicy() { return dbGet("SELECT * FROM identity_policies WHERE id = 'default'") || {}; } function ensureSsoEnabled() { if (!identityPolicy().sso_enabled) throw oidcError(403, "sso_disabled", "平台当前未启用企业 SSO"); } function cleanupExpired() { const timestamp = nowIso(); dbRun("DELETE FROM oidc_login_states WHERE expires_at <= ? OR consumed_at IS NOT NULL", [timestamp]); dbRun("DELETE FROM auth_sso_tickets WHERE expires_at <= ? OR consumed_at IS NOT NULL", [timestamp]); } async function fetchJson(url, options = {}, label = "OIDC 请求") { const controller = new AbortController(); const timeout = setTimeout(() => controller.abort(), 8000); try { const response = await fetch(url, { ...options, signal: controller.signal }); const raw = await response.text(); const payload = raw ? parseJson(raw, {}) : {}; if (!response.ok) throw oidcError(502, "oidc_upstream_error", `${label}失败:HTTP ${response.status}`, { status: response.status, response: payload }); return payload; } catch (error) { if (error.status) throw error; throw oidcError(502, "oidc_upstream_unreachable", `${label}失败:${error.message}`); } finally { clearTimeout(timeout); } } async function discover(provider) { const discoveryUrl = `${normalizeIssuer(provider.issuer_url)}/.well-known/openid-configuration`; const discovery = await fetchJson(discoveryUrl, { headers: { accept: "application/json" } }, "OIDC discovery"); return { issuer: discovery.issuer || provider.issuer_url, authorizationEndpoint: provider.authorization_url || discovery.authorization_endpoint, tokenEndpoint: provider.token_url || discovery.token_endpoint, userinfoEndpoint: provider.userinfo_url || discovery.userinfo_endpoint || "", jwksUri: provider.jwks_url || discovery.jwks_uri || "" }; } export function startOidcLogin(providerId, { redirectUri, returnTo = "/", selection = {}, ipAddress = "", userAgent = "" } = {}) { ensureSsoEnabled(); const provider = providerRow(providerId); const authorizationUrl = String(provider.authorization_url || "").trim(); if (!authorizationUrl) throw oidcError(400, "sso_authorization_endpoint_missing", "OIDC 提供商尚未配置 Authorization Endpoint,请先探测"); const state = base64url(randomBytes(32)); const nonce = base64url(randomBytes(32)); const codeVerifier = base64url(randomBytes(48)); const codeChallenge = base64url(createHash("sha256").update(codeVerifier).digest()); const timestamp = nowIso(); const expiresAt = new Date(Date.now() + LOGIN_STATE_TTL_MS).toISOString(); const id = `oidc-state-${Date.now()}-${randomBytes(4).toString("hex")}`; dbRun("INSERT INTO oidc_login_states(id, state_hash, provider_id, nonce_hash, code_verifier_ciphertext, redirect_uri, return_to, selection_json, ip_address, user_agent, expires_at, created_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)", [ id, hashValue(state), provider.id, hashValue(nonce), encryptOpaque(codeVerifier), redirectUri, String(returnTo || "/"), JSON.stringify(selection || {}), String(ipAddress || ""), String(userAgent || ""), expiresAt, timestamp ]); const url = new URL(authorizationUrl); const scopes = parseJson(provider.scopes_json, ["openid", "profile", "email"]); url.searchParams.set("response_type", "code"); url.searchParams.set("client_id", provider.client_id); url.searchParams.set("redirect_uri", redirectUri); url.searchParams.set("scope", Array.isArray(scopes) ? scopes.join(" ") : "openid profile email"); url.searchParams.set("state", state); url.searchParams.set("nonce", nonce); url.searchParams.set("code_challenge", codeChallenge); url.searchParams.set("code_challenge_method", "S256"); return { provider: { id: provider.id, name: provider.name }, authorizationUrl: url.toString(), expiresAt }; } function consumeLoginState(state) { cleanupExpired(); const timestamp = nowIso(); const result = dbRun("UPDATE oidc_login_states SET consumed_at = ? WHERE state_hash = ? AND consumed_at IS NULL AND expires_at > ?", [timestamp, hashValue(state), timestamp]); if (!Number(result.changes || 0)) throw oidcError(401, "oidc_state_invalid", "OIDC 登录状态无效、已使用或已过期"); const row = dbGet("SELECT * FROM oidc_login_states WHERE state_hash = ?", [hashValue(state)]); if (!row) throw oidcError(401, "oidc_state_invalid", "OIDC 登录状态不存在"); return row; } function parseJwt(token) { const parts = String(token || "").split("."); if (parts.length !== 3) throw oidcError(502, "oidc_id_token_invalid", "OIDC 返回的 ID Token 格式无效"); try { return { header: JSON.parse(decodeBase64url(parts[0]).toString("utf8")), claims: JSON.parse(decodeBase64url(parts[1]).toString("utf8")), encodedHeader: parts[0], encodedPayload: parts[1], signature: decodeBase64url(parts[2]) }; } catch { throw oidcError(502, "oidc_id_token_invalid", "OIDC 返回的 ID Token 不是有效 JWT"); } } function ecdsaJoseToDer(signature) { const half = Math.floor(signature.length / 2); const encodeInteger = (value) => { let output = Buffer.from(value); while (output.length > 1 && output[0] === 0) output = output.subarray(1); if (output[0] & 0x80) output = Buffer.concat([Buffer.from([0]), output]); return Buffer.concat([Buffer.from([0x02, output.length]), output]); }; const sequence = Buffer.concat([encodeInteger(signature.subarray(0, half)), encodeInteger(signature.subarray(half))]); if (sequence.length >= 128) return Buffer.concat([Buffer.from([0x30, 0x81, sequence.length]), sequence]); return Buffer.concat([Buffer.from([0x30, sequence.length]), sequence]); } async function verifyIdToken(token, provider, discovery) { const parsed = parseJwt(token); const { header, claims, encodedHeader, encodedPayload, signature } = parsed; const algorithms = { RS256: { hash: "SHA256", type: "rsa" }, RS384: { hash: "SHA384", type: "rsa" }, RS512: { hash: "SHA512", type: "rsa" }, PS256: { hash: "SHA256", type: "pss" }, PS384: { hash: "SHA384", type: "pss" }, PS512: { hash: "SHA512", type: "pss" }, ES256: { hash: "SHA256", type: "ecdsa" }, ES384: { hash: "SHA384", type: "ecdsa" }, ES512: { hash: "SHA512", type: "ecdsa" } }; const algorithm = algorithms[header.alg]; if (!algorithm) throw oidcError(502, "oidc_algorithm_unsupported", `OIDC ID Token 签名算法不受支持:${header.alg || "未声明"}`); if (!discovery.jwksUri) throw oidcError(502, "oidc_jwks_missing", "OIDC 提供商未返回 JWKS 地址"); let cache = OIDC_JWKS_CACHE.get(discovery.jwksUri); if (!cache || cache.expiresAt <= Date.now()) { cache = { keys: (await fetchJson(discovery.jwksUri, { headers: { accept: "application/json" } }, "OIDC JWKS")).keys || [], expiresAt: Date.now() + 5 * 60 * 1000 }; OIDC_JWKS_CACHE.set(discovery.jwksUri, cache); } let jwk = cache.keys.find((item) => item.kid && item.kid === header.kid); if (!jwk && !header.kid && cache.keys.length === 1) jwk = cache.keys[0]; if (!jwk) { OIDC_JWKS_CACHE.delete(discovery.jwksUri); const refreshed = (await fetchJson(discovery.jwksUri, { headers: { accept: "application/json" } }, "OIDC JWKS 刷新")).keys || []; OIDC_JWKS_CACHE.set(discovery.jwksUri, { keys: refreshed, expiresAt: Date.now() + 5 * 60 * 1000 }); jwk = refreshed.find((item) => item.kid && item.kid === header.kid) || (!header.kid && refreshed.length === 1 ? refreshed[0] : null); } if (!jwk) throw oidcError(502, "oidc_signing_key_not_found", "OIDC ID Token 的签名密钥不在 JWKS 中"); let publicKey; try { publicKey = createPublicKey({ key: jwk, format: "jwk" }); } catch (error) { throw oidcError(502, "oidc_jwk_invalid", `OIDC JWKS 公钥无效:${error.message}`); } const verify = createVerify(algorithm.hash); verify.update(`${encodedHeader}.${encodedPayload}`); verify.end(); const normalizedSignature = algorithm.type === "ecdsa" ? ecdsaJoseToDer(signature) : signature; const verified = algorithm.type === "pss" ? verify.verify({ key: publicKey, padding: RSA_PKCS1_PSS_PADDING, saltLength: RSA_PSS_SALTLEN_DIGEST }, normalizedSignature) : verify.verify(publicKey, normalizedSignature); if (!verified) throw oidcError(401, "oidc_signature_invalid", "OIDC ID Token 签名校验失败"); const now = Date.now(); if (normalizeIssuer(claims.iss) !== normalizeIssuer(provider.issuer_url)) throw oidcError(401, "oidc_issuer_invalid", "OIDC ID Token 的 Issuer 不匹配"); const audiences = Array.isArray(claims.aud) ? claims.aud : [claims.aud]; if (!audiences.includes(provider.client_id)) throw oidcError(401, "oidc_audience_invalid", "OIDC ID Token 的 Audience 不匹配"); if (claims.azp && claims.azp !== provider.client_id) throw oidcError(401, "oidc_authorized_party_invalid", "OIDC ID Token 的 azp 不匹配"); if (!claims.nonce) throw oidcError(401, "oidc_nonce_missing", "OIDC ID Token 缺少 nonce"); if (!claims.sub) throw oidcError(401, "oidc_subject_missing", "OIDC ID Token 缺少 subject"); if (!claims.exp || Number(claims.exp) * 1000 + CLOCK_SKEW_MS < now) throw oidcError(401, "oidc_token_expired", "OIDC ID Token 已过期"); if (claims.iat && Number(claims.iat) * 1000 - CLOCK_SKEW_MS > now) throw oidcError(401, "oidc_token_issued_in_future", "OIDC ID Token 的签发时间无效"); return claims; } async function exchangeCode(provider, stateRow, code) { const discovery = await discover(provider); if (!discovery.tokenEndpoint) throw oidcError(400, "sso_token_endpoint_missing", "OIDC 提供商尚未配置 Token Endpoint"); const clientSecret = process.env[provider.client_secret_ref]; if (!clientSecret) throw oidcError(503, "sso_client_secret_missing", `服务端环境变量 ${provider.client_secret_ref} 未配置`); const body = new URLSearchParams({ grant_type: "authorization_code", code: String(code || ""), redirect_uri: stateRow.redirect_uri, client_id: provider.client_id, client_secret: clientSecret, code_verifier: decryptOpaque(stateRow.code_verifier_ciphertext) }); const token = await fetchJson(discovery.tokenEndpoint, { method: "POST", headers: { "content-type": "application/x-www-form-urlencoded", accept: "application/json" }, body }, "OIDC code exchange"); if (!token.id_token) throw oidcError(502, "oidc_id_token_missing", "OIDC Token Endpoint 未返回 ID Token"); const claims = await verifyIdToken(token.id_token, provider, discovery); if (hashValue(claims.nonce) !== stateRow.nonce_hash) throw oidcError(401, "oidc_nonce_invalid", "OIDC ID Token 的 nonce 不匹配"); if (discovery.userinfoEndpoint && token.access_token) { try { const userInfo = await fetchJson(discovery.userinfoEndpoint, { headers: { authorization: `Bearer ${token.access_token}`, accept: "application/json" } }, "OIDC UserInfo"); for (const [key, value] of Object.entries(userInfo || {})) if (claims[key] === undefined) claims[key] = value; } catch { // ID Token claims remain authoritative when an optional UserInfo call fails. } } return { claims, token, discovery }; } function avatarColorFor(email) { const colors = ["#d97757", "#3f7f87", "#a77646", "#8e6a9f", "#477d69", "#9a6b51"]; const digest = createHash("sha256").update(email).digest().readUInt16BE(0); return colors[digest % colors.length]; } function organizationFor(provider, existingUser) { if (provider.organization_id) { const organization = dbGet("SELECT * FROM organizations WHERE id = ? AND status = 'active'", [provider.organization_id]); if (!organization) throw oidcError(403, "sso_organization_unavailable", "SSO 提供商绑定的组织不存在或已停用"); return organization; } if (existingUser) { const membership = dbGet("SELECT o.* FROM organizations o JOIN organization_members om ON om.organization_id = o.id WHERE om.user_id = ? AND om.status = 'active' AND o.status = 'active' ORDER BY om.joined_at ASC LIMIT 1", [existingUser.id]); if (membership) return membership; } throw oidcError(403, "sso_organization_required", "该 SSO 提供商尚未绑定组织,不能自动创建企业成员"); } function ensureOrganizationMembership(user, organization, provider) { const current = dbGet("SELECT * FROM organization_members WHERE organization_id = ? AND user_id = ?", [organization.id, user.id]); if (current?.status === "active") return current; if (current && current.status !== "active" && !provider.auto_provision) throw oidcError(403, "sso_membership_inactive", "当前用户在目标组织中已被停用"); if (!current && !provider.auto_provision) throw oidcError(403, "sso_membership_required", "当前企业账号尚未加入目标组织"); const timestamp = nowIso(); const roleKey = provider.default_role_key || "org_member"; dbRun("INSERT INTO organization_members(id, organization_id, user_id, role_key, status, joined_at, created_at, updated_at) VALUES (?, ?, ?, ?, 'active', ?, ?, ?) ON CONFLICT(organization_id, user_id) DO UPDATE SET role_key = excluded.role_key, status = 'active', joined_at = excluded.joined_at, updated_at = excluded.updated_at", [`om-sso-${organization.id}-${user.id}`, organization.id, user.id, roleKey, timestamp, timestamp, timestamp]); return dbGet("SELECT * FROM organization_members WHERE organization_id = ? AND user_id = ?", [organization.id, user.id]); } function ensureWorkspaceMembership(user, organization, provider) { const workspace = provider.workspace_id ? dbGet("SELECT * FROM workspaces WHERE id = ? AND organization_id = ? AND status = 'active'", [provider.workspace_id, organization.id]) : dbGet("SELECT * FROM workspaces WHERE organization_id = ? AND status = 'active' ORDER BY created_at ASC LIMIT 1", [organization.id]); if (!workspace) throw oidcError(403, "sso_workspace_required", "SSO 提供商没有可用的默认工作区"); const current = dbGet("SELECT * FROM workspace_members WHERE workspace_id = ? AND user_id = ?", [workspace.id, user.id]); if (current?.status === "active") return { workspace, membership: current }; if (current && current.status !== "active" && !provider.auto_provision) throw oidcError(403, "sso_workspace_membership_inactive", "当前用户在默认工作区中已被停用"); if (!current && !provider.auto_provision) throw oidcError(403, "sso_workspace_membership_required", "当前企业账号尚未加入默认工作区"); const timestamp = nowIso(); const roleKey = provider.default_workspace_role_key || "writer"; dbRun("INSERT INTO workspace_members(id, workspace_id, user_id, role_key, status, created_at, updated_at) VALUES (?, ?, ?, ?, 'active', ?, ?) ON CONFLICT(workspace_id, user_id) DO UPDATE SET role_key = excluded.role_key, status = 'active', updated_at = excluded.updated_at", [`wm-sso-${workspace.id}-${user.id}`, workspace.id, user.id, roleKey, timestamp, timestamp]); return { workspace, membership: dbGet("SELECT * FROM workspace_members WHERE workspace_id = ? AND user_id = ?", [workspace.id, user.id]) }; } export function resolveOidcUser(provider, claims) { const mapping = parseJson(provider.claim_mapping_json, { email: "email", displayName: "name", externalId: "sub" }); const subject = String(claimValue(claims, mapping.externalId || "sub") || claims.sub || "").trim(); const email = String(claimValue(claims, mapping.email || "email") || claims.preferred_username || "").trim().toLowerCase(); const displayName = String(claimValue(claims, mapping.displayName || "name") || claims.preferred_username || email || subject).trim(); if (!subject) throw oidcError(401, "sso_subject_missing", "企业身份没有返回可用的 subject"); if (!email || !email.includes("@")) throw oidcError(401, "sso_email_missing", "企业身份没有返回可用邮箱,无法完成平台账号绑定"); const existingIdentity = dbGet("SELECT ei.*, u.* FROM external_identities ei JOIN users u ON u.id = ei.user_id WHERE ei.provider_id = ? AND ei.subject = ?", [provider.id, subject]); let user = existingIdentity ? dbGet("SELECT * FROM users WHERE id = ?", [existingIdentity.user_id]) : dbGet("SELECT * FROM users WHERE lower(email) = ?", [email]); const organization = organizationFor(provider, user); if (existingIdentity && existingIdentity.email_at_login && existingIdentity.email_at_login !== email && user?.email !== email) throw oidcError(403, "sso_identity_mismatch", "企业身份 subject 与邮箱绑定不一致,需要管理员处理"); if (!user) { if (!provider.auto_provision) throw oidcError(403, "sso_auto_provision_disabled", "该企业账号尚未注册,且当前 SSO 提供商关闭了自动入组"); const timestamp = nowIso(); const id = `u-sso-${Date.now()}-${randomBytes(5).toString("hex")}`; dbRun("INSERT INTO users(id, display_name, email, avatar_color, status, created_at, updated_at) VALUES (?, ?, ?, ?, 'active', ?, ?)", [id, displayName.slice(0, 120) || email, email, avatarColorFor(email), timestamp, timestamp]); user = dbGet("SELECT * FROM users WHERE id = ?", [id]); } if (!user || user.status !== "active") throw oidcError(403, "user_not_active", "当前企业账号已停用"); ensureOrganizationMembership(user, organization, provider); ensureWorkspaceMembership(user, organization, provider); const timestamp = nowIso(); if (!existingIdentity) { const identityId = `ext-${Date.now()}-${randomBytes(5).toString("hex")}`; try { dbRun("INSERT INTO external_identities(id, provider_id, user_id, subject, issuer, email_at_login, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?)", [identityId, provider.id, user.id, subject, normalizeIssuer(claims.iss || provider.issuer_url), email, timestamp, timestamp]); } catch (error) { const conflict = dbGet("SELECT user_id FROM external_identities WHERE provider_id = ? AND subject = ?", [provider.id, subject]); if (!conflict || conflict.user_id !== user.id) throw oidcError(409, "sso_identity_already_bound", "该企业身份已经绑定其他平台用户"); } } else { dbRun("UPDATE external_identities SET email_at_login = ?, issuer = ?, updated_at = ? WHERE id = ?", [email, normalizeIssuer(claims.iss || provider.issuer_url), timestamp, existingIdentity.id]); } return { user: dbGet("SELECT * FROM users WHERE id = ?", [user.id]), organization, subject, email }; } export async function handleOidcCallback({ code, state }) { const stateRow = consumeLoginState(state); const provider = providerRow(stateRow.provider_id); const exchanged = await exchangeCode(provider, stateRow, code); return { ...resolveOidcUser(provider, exchanged.claims), selection: parseJson(stateRow.selection_json, {}), returnTo: stateRow.return_to, provider }; } export function createSsoTicket(userId, selection = {}, metadata = {}) { cleanupExpired(); const ticket = base64url(randomBytes(32)); const timestamp = nowIso(); const expiresAt = new Date(Date.now() + TICKET_TTL_MS).toISOString(); const id = `sso-ticket-${Date.now()}-${randomBytes(4).toString("hex")}`; dbRun("INSERT INTO auth_sso_tickets(id, ticket_hash, user_id, selection_json, ip_address, user_agent, expires_at, created_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?)", [id, hashValue(ticket), userId, JSON.stringify(selection || {}), String(metadata.ipAddress || ""), String(metadata.userAgent || ""), expiresAt, timestamp]); return { ticket, expiresAt }; } export function redeemSsoTicket(ticket, metadata = {}) { cleanupExpired(); const timestamp = nowIso(); const result = dbRun("UPDATE auth_sso_tickets SET consumed_at = ? WHERE ticket_hash = ? AND consumed_at IS NULL AND expires_at > ?", [timestamp, hashValue(ticket), timestamp]); if (!Number(result.changes || 0)) throw oidcError(401, "sso_ticket_invalid", "SSO 登录票据无效、已使用或已过期"); const row = dbGet("SELECT * FROM auth_sso_tickets WHERE ticket_hash = ?", [hashValue(ticket)]); if (!row) throw oidcError(401, "sso_ticket_invalid", "SSO 登录票据不存在"); const user = dbGet("SELECT * FROM users WHERE id = ?", [row.user_id]); if (!user || user.status !== "active") throw oidcError(403, "user_not_active", "当前企业账号已停用"); const selection = parseJson(row.selection_json, {}); if (mfaStatus(user.id).enabled) return { mfaRequired: true, challenge: createMfaChallenge(user.id), user: safeUser(user), selection }; if (mfaRequiredForUser(user.id)) return { mfaRequired: true, mfaEnrollmentRequired: true, enrollment: createMfaEnrollmentChallenge(user.id), user: safeUser(user), selection }; return { session: createSession(user.id, metadata), user: safeUser(user), selection }; }