Refactor validate jws method

This commit is contained in:
Fang-Pen Lin
2025-11-07 09:18:15 -08:00
parent 4f0bb7fb8a
commit 540521588d
4 changed files with 55 additions and 46 deletions
@@ -61,9 +61,9 @@ export const pkiAcmeAccountDALFactory = (db: TDbClient) => {
} }
}; };
const findByPublicKey = async (publicKey: unknown, alg: string, tx?: Knex) => { const findByPublicKey = async (profileId: string, alg: string, publicKey: unknown, tx?: Knex) => {
try { try {
const account = await (tx || db)(TableName.PkiAcmeAccount).where({ publicKey, alg }).first(); const account = await (tx || db)(TableName.PkiAcmeAccount).where({ profileId, alg, publicKey }).first();
return account || null; return account || null;
} catch (error) { } catch (error) {
@@ -3,14 +3,19 @@ import { NotFoundError } from "@app/lib/errors";
import { TCertificateProfileDALFactory } from "@app/services/certificate-profile/certificate-profile-dal"; import { TCertificateProfileDALFactory } from "@app/services/certificate-profile/certificate-profile-dal";
import { AcmeAccountDoesNotExistError, AcmeBadPublicKeyError, AcmeMalformedError } from "./pki-acme-errors"; import {
AcmeAccountDoesNotExistError,
AcmeBadPublicKeyError,
AcmeMalformedError,
AcmeServerInternalError
} from "./pki-acme-errors";
import { TPkiAcmeAccounts } from "@app/db/schemas/pki-acme-accounts"; import { TPkiAcmeAccounts } from "@app/db/schemas/pki-acme-accounts";
import { import {
EnrollmentType, EnrollmentType,
TCertificateProfileWithConfigs TCertificateProfileWithConfigs
} from "@app/services/certificate-profile/certificate-profile-types"; } from "@app/services/certificate-profile/certificate-profile-types";
import { flattenedVerify, importJWK, JWK, JWSHeaderParameters } from "jose"; import { flattenedVerify, FlattenedVerifyResult, importJWK, JWK, JWSHeaderParameters } from "jose";
import { TPkiAcmeAccountDALFactory } from "./pki-acme-account-dal"; import { TPkiAcmeAccountDALFactory } from "./pki-acme-account-dal";
import { TPkiAcmeOrderDALFactory } from "./pki-acme-order-dal"; import { TPkiAcmeOrderDALFactory } from "./pki-acme-order-dal";
import { ProtectedHeaderSchema } from "./pki-acme-schemas"; import { ProtectedHeaderSchema } from "./pki-acme-schemas";
@@ -27,12 +32,16 @@ import {
TGetAcmeAuthorizationResponse, TGetAcmeAuthorizationResponse,
TGetAcmeDirectoryResponse, TGetAcmeDirectoryResponse,
TGetAcmeOrderResponse, TGetAcmeOrderResponse,
TJwsPayload,
TJwsPayloadWithJwk, TJwsPayloadWithJwk,
TListAcmeOrdersResponse, TListAcmeOrdersResponse,
TPkiAcmeServiceFactory, TPkiAcmeServiceFactory,
TRawJwsPayload, TRawJwsPayload,
TRespondToAcmeChallengeResponse TRespondToAcmeChallengeResponse
} from "./pki-acme-types"; } from "./pki-acme-types";
import { JWSInvalid } from "jose/dist/types/util/errors";
import { logger } from "@app/lib/logger";
import { z, ZodError } from "zod";
type TPkiAcmeServiceFactoryDep = { type TPkiAcmeServiceFactoryDep = {
certificateProfileDAL: Pick<TCertificateProfileDALFactory, "findById">; certificateProfileDAL: Pick<TCertificateProfileDALFactory, "findById">;
@@ -61,57 +70,52 @@ export const pkiAcmeServiceFactory = ({
return `${baseUrl}${path}`; return `${baseUrl}${path}`;
}; };
const validateCreateAcmeAccountJwsPayload = async (rawJwsPayload: TRawJwsPayload): Promise<TJwsPayloadWithJwk> => { const validateJwsPayload = async <T>(
const { payload: rawPayload, protectedHeader: rawProtectedHeader } = await flattenedVerify( rawJwsPayload: TRawJwsPayload,
rawJwsPayload, getJWK: (protectedHeader: JWSHeaderParameters) => Promise<JsonWebKey>,
async (protectedHeader: JWSHeaderParameters | undefined) => { schema: z.ZodSchema<T>
): Promise<TJwsPayload> => {
let result: FlattenedVerifyResult;
try {
result = await flattenedVerify(rawJwsPayload, async (protectedHeader: JWSHeaderParameters | undefined) => {
if (protectedHeader === undefined) { if (protectedHeader === undefined) {
throw new AcmeMalformedError({ detail: "Protected header is required" }); throw new AcmeMalformedError({ detail: "Protected header is required" });
} }
if (protectedHeader.jwk === undefined) { ProtectedHeaderSchema.parse(protectedHeader);
throw new AcmeBadPublicKeyError({ detail: "JWK is required in the protected header" }); const jwk = await getJWK(protectedHeader);
return await importJWK(jwk, protectedHeader.alg);
});
} catch (error) {
if (error instanceof ZodError) {
throw new AcmeMalformedError({ detail: `Invalid JWS payload: ${error.message}` });
} }
// For the create account request, the JWK is provided in the protected header. if (error instanceof JWSInvalid) {
// Let use it to verify the signature. throw new AcmeBadPublicKeyError({ detail: "Invalid JWS payload" });
const imported = await importJWK(protectedHeader.jwk as JWK, protectedHeader.alg);
return imported;
} }
); logger.error(error, "Unexpected error while verifying JWS payload");
throw new AcmeServerInternalError({ detail: "Failed to verify JWS payload" });
}
const { payload: rawPayload, protectedHeader: rawProtectedHeader } = result!;
const { success, data: protectedHeader } = ProtectedHeaderSchema.safeParse(rawProtectedHeader); const { success, data: protectedHeader } = ProtectedHeaderSchema.safeParse(rawProtectedHeader);
if (!success) { if (!success) {
throw new AcmeMalformedError({ detail: "Invalid protected header" }); throw new AcmeMalformedError({ detail: "Invalid protected header" });
} }
const decoder = new TextDecoder(); const decoder = new TextDecoder();
const payload = JSON.parse(decoder.decode(rawPayload)) as TCreateAcmeAccountPayload; const jsonPayload = JSON.parse(decoder.decode(rawPayload));
// TODO: also consume the nonce here try {
return { payload, protectedHeader, jwk: protectedHeader.jwk as JsonWebKey, alg: protectedHeader.alg }; const payload = schema.parse(jsonPayload);
return {
protectedHeader,
payload
}; };
} catch (error) {
const validateJwsPayload = async (rawJwsPayload: TRawJwsPayload): Promise<TJwsPayloadWithJwk> => { if (error instanceof ZodError) {
const { payload: rawPayload, protectedHeader: rawProtectedHeader } = await flattenedVerify( throw new AcmeMalformedError({ detail: `Invalid JWS payload: ${error.message}` });
rawJwsPayload,
async (protectedHeader: JWSHeaderParameters | undefined) => {
if (protectedHeader === undefined) {
throw new AcmeMalformedError({ detail: "Protected header is required" });
} }
if (protectedHeader.kid === undefined) { logger.error(error, "Unexpected error while parsing JWS payload");
throw new AcmeBadPublicKeyError({ detail: "Kid is required in the protected header" }); throw new AcmeServerInternalError({ detail: "Failed to verify JWS payload" });
} }
const imported = await importJWK(protectedHeader.jwk as JWK, protectedHeader.alg);
return imported;
}
);
const { success, data: protectedHeader } = ProtectedHeaderSchema.safeParse(rawProtectedHeader);
if (!success) {
throw new AcmeMalformedError({ detail: "Invalid protected header" });
}
const decoder = new TextDecoder();
const payload = JSON.parse(decoder.decode(rawPayload)) as TCreateAcmeAccountPayload;
// TODO: also consume the nonce here
return { payload, protectedHeader, jwk: protectedHeader.jwk as JsonWebKey };
}; };
const getAcmeDirectory = async (profileId: string): Promise<TGetAcmeDirectoryResponse> => { const getAcmeDirectory = async (profileId: string): Promise<TGetAcmeDirectoryResponse> => {
@@ -132,11 +136,12 @@ export const pkiAcmeServiceFactory = ({
const createAcmeAccount = async ( const createAcmeAccount = async (
profileId: string, profileId: string,
alg: string,
jwk: JWK, jwk: JWK,
{ onlyReturnExisting, contact }: TCreateAcmeAccountPayload { onlyReturnExisting, contact }: TCreateAcmeAccountPayload
): Promise<TAcmeResponse<TCreateAcmeAccountResponse>> => { ): Promise<TAcmeResponse<TCreateAcmeAccountResponse>> => {
const profile = await validateAcmeProfile(profileId); const profile = await validateAcmeProfile(profileId);
const existingAccount: TPkiAcmeAccounts | null = await acmeAccountDAL.findByPublicKey(jwk, alg); const existingAccount: TPkiAcmeAccounts | null = await acmeAccountDAL.findByPublicKey(profileId, alg, jwk);
if (onlyReturnExisting && !existingAccount) { if (onlyReturnExisting && !existingAccount) {
throw new AcmeAccountDoesNotExistError({ message: "ACME account not found" }); throw new AcmeAccountDoesNotExistError({ message: "ACME account not found" });
} }
@@ -55,6 +55,7 @@ export type TPkiAcmeServiceFactory = {
getAcmeNewNonce: (profileId: string) => Promise<string>; getAcmeNewNonce: (profileId: string) => Promise<string>;
createAcmeAccount: ( createAcmeAccount: (
profileId: string, profileId: string,
alg: string,
jwk: JsonWebKey, jwk: JsonWebKey,
body: TCreateAcmeAccountPayload body: TCreateAcmeAccountPayload
) => Promise<TAcmeResponse<TCreateAcmeAccountResponse>>; ) => Promise<TAcmeResponse<TCreateAcmeAccountResponse>>;
+5 -2
View File
@@ -361,6 +361,7 @@ import { initializeOauthConfigSync } from "./v1/sso-router";
import { registerV2Routes } from "./v2"; import { registerV2Routes } from "./v2";
import { registerV3Routes } from "./v3"; import { registerV3Routes } from "./v3";
import { registerV4Routes } from "./v4"; import { registerV4Routes } from "./v4";
import { pkiAcmeOrderDALFactory } from "@app/ee/services/pki-acme/pki-acme-order-dal";
const histogram = monitorEventLoopDelay({ resolution: 20 }); const histogram = monitorEventLoopDelay({ resolution: 20 });
histogram.enable(); histogram.enable();
@@ -1065,7 +1066,8 @@ export const registerRoutes = async (
const apiEnrollmentConfigDAL = apiEnrollmentConfigDALFactory(db); const apiEnrollmentConfigDAL = apiEnrollmentConfigDALFactory(db);
const estEnrollmentConfigDAL = estEnrollmentConfigDALFactory(db); const estEnrollmentConfigDAL = estEnrollmentConfigDALFactory(db);
const acmeEnrollmentConfigDAL = acmeEnrollmentConfigDALFactory(db); const acmeEnrollmentConfigDAL = acmeEnrollmentConfigDALFactory(db);
const pkiAcmeAccountDAL = pkiAcmeAccountDALFactory(db); const acmeAccountDAL = pkiAcmeAccountDALFactory(db);
const acmeOrderDAL = pkiAcmeOrderDALFactory(db);
const certificateDAL = certificateDALFactory(db); const certificateDAL = certificateDALFactory(db);
const certificateBodyDAL = certificateBodyDALFactory(db); const certificateBodyDAL = certificateBodyDALFactory(db);
@@ -1169,7 +1171,8 @@ export const registerRoutes = async (
const pkiAcmeService = pkiAcmeServiceFactory({ const pkiAcmeService = pkiAcmeServiceFactory({
certificateProfileDAL, certificateProfileDAL,
pkiAcmeAccountDAL acmeAccountDAL,
acmeOrderDAL
}); });
const pkiAlertService = pkiAlertServiceFactory({ const pkiAlertService = pkiAlertServiceFactory({