Refactor validate jws method

This commit is contained in:
Fang-Pen Lin
2025-10-28 19:02:17 -07:00
parent 4f0bb7fb8a
commit 540521588d
4 changed files with 55 additions and 46 deletions

View File

@@ -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 {
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;
} catch (error) {

View File

@@ -3,14 +3,19 @@ import { NotFoundError } from "@app/lib/errors";
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 {
EnrollmentType,
TCertificateProfileWithConfigs
} 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 { TPkiAcmeOrderDALFactory } from "./pki-acme-order-dal";
import { ProtectedHeaderSchema } from "./pki-acme-schemas";
@@ -27,12 +32,16 @@ import {
TGetAcmeAuthorizationResponse,
TGetAcmeDirectoryResponse,
TGetAcmeOrderResponse,
TJwsPayload,
TJwsPayloadWithJwk,
TListAcmeOrdersResponse,
TPkiAcmeServiceFactory,
TRawJwsPayload,
TRespondToAcmeChallengeResponse
} 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 = {
certificateProfileDAL: Pick<TCertificateProfileDALFactory, "findById">;
@@ -61,57 +70,52 @@ export const pkiAcmeServiceFactory = ({
return `${baseUrl}${path}`;
};
const validateCreateAcmeAccountJwsPayload = async (rawJwsPayload: TRawJwsPayload): Promise<TJwsPayloadWithJwk> => {
const { payload: rawPayload, protectedHeader: rawProtectedHeader } = await flattenedVerify(
rawJwsPayload,
async (protectedHeader: JWSHeaderParameters | undefined) => {
const validateJwsPayload = async <T>(
rawJwsPayload: TRawJwsPayload,
getJWK: (protectedHeader: JWSHeaderParameters) => Promise<JsonWebKey>,
schema: z.ZodSchema<T>
): Promise<TJwsPayload> => {
let result: FlattenedVerifyResult;
try {
result = await flattenedVerify(rawJwsPayload, async (protectedHeader: JWSHeaderParameters | undefined) => {
if (protectedHeader === undefined) {
throw new AcmeMalformedError({ detail: "Protected header is required" });
}
if (protectedHeader.jwk === undefined) {
throw new AcmeBadPublicKeyError({ detail: "JWK is required in the protected header" });
}
// For the create account request, the JWK is provided in the protected header.
// Let use it to verify the signature.
const imported = await importJWK(protectedHeader.jwk as JWK, protectedHeader.alg);
return imported;
ProtectedHeaderSchema.parse(protectedHeader);
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}` });
}
);
if (error instanceof JWSInvalid) {
throw new AcmeBadPublicKeyError({ detail: "Invalid JWS payload" });
}
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);
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, alg: protectedHeader.alg };
};
const validateJwsPayload = async (rawJwsPayload: TRawJwsPayload): Promise<TJwsPayloadWithJwk> => {
const { payload: rawPayload, protectedHeader: rawProtectedHeader } = await flattenedVerify(
rawJwsPayload,
async (protectedHeader: JWSHeaderParameters | undefined) => {
if (protectedHeader === undefined) {
throw new AcmeMalformedError({ detail: "Protected header is required" });
}
if (protectedHeader.kid === undefined) {
throw new AcmeBadPublicKeyError({ detail: "Kid is required in the protected header" });
}
const imported = await importJWK(protectedHeader.jwk as JWK, protectedHeader.alg);
return imported;
const jsonPayload = JSON.parse(decoder.decode(rawPayload));
try {
const payload = schema.parse(jsonPayload);
return {
protectedHeader,
payload
};
} catch (error) {
if (error instanceof ZodError) {
throw new AcmeMalformedError({ detail: `Invalid JWS payload: ${error.message}` });
}
);
const { success, data: protectedHeader } = ProtectedHeaderSchema.safeParse(rawProtectedHeader);
if (!success) {
throw new AcmeMalformedError({ detail: "Invalid protected header" });
logger.error(error, "Unexpected error while parsing JWS payload");
throw new AcmeServerInternalError({ detail: "Failed to verify JWS payload" });
}
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> => {
@@ -132,11 +136,12 @@ export const pkiAcmeServiceFactory = ({
const createAcmeAccount = async (
profileId: string,
alg: string,
jwk: JWK,
{ onlyReturnExisting, contact }: TCreateAcmeAccountPayload
): Promise<TAcmeResponse<TCreateAcmeAccountResponse>> => {
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) {
throw new AcmeAccountDoesNotExistError({ message: "ACME account not found" });
}

View File

@@ -55,6 +55,7 @@ export type TPkiAcmeServiceFactory = {
getAcmeNewNonce: (profileId: string) => Promise<string>;
createAcmeAccount: (
profileId: string,
alg: string,
jwk: JsonWebKey,
body: TCreateAcmeAccountPayload
) => Promise<TAcmeResponse<TCreateAcmeAccountResponse>>;

View File

@@ -361,6 +361,7 @@ import { initializeOauthConfigSync } from "./v1/sso-router";
import { registerV2Routes } from "./v2";
import { registerV3Routes } from "./v3";
import { registerV4Routes } from "./v4";
import { pkiAcmeOrderDALFactory } from "@app/ee/services/pki-acme/pki-acme-order-dal";
const histogram = monitorEventLoopDelay({ resolution: 20 });
histogram.enable();
@@ -1065,7 +1066,8 @@ export const registerRoutes = async (
const apiEnrollmentConfigDAL = apiEnrollmentConfigDALFactory(db);
const estEnrollmentConfigDAL = estEnrollmentConfigDALFactory(db);
const acmeEnrollmentConfigDAL = acmeEnrollmentConfigDALFactory(db);
const pkiAcmeAccountDAL = pkiAcmeAccountDALFactory(db);
const acmeAccountDAL = pkiAcmeAccountDALFactory(db);
const acmeOrderDAL = pkiAcmeOrderDALFactory(db);
const certificateDAL = certificateDALFactory(db);
const certificateBodyDAL = certificateBodyDALFactory(db);
@@ -1169,7 +1171,8 @@ export const registerRoutes = async (
const pkiAcmeService = pkiAcmeServiceFactory({
certificateProfileDAL,
pkiAcmeAccountDAL
acmeAccountDAL,
acmeOrderDAL
});
const pkiAlertService = pkiAlertServiceFactory({