feat: bootstrap commercial AI drama platform
This commit is contained in:
+397
@@ -0,0 +1,397 @@
|
||||
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 };
|
||||
}
|
||||
Reference in New Issue
Block a user