diff --git a/lsquic/context/context.nim b/lsquic/context/context.nim index fed33f0..62f7f82 100644 --- a/lsquic/context/context.nim +++ b/lsquic/context/context.nim @@ -247,21 +247,32 @@ proc alpnSelectProtoCB( proc verifyCertificate( ssl: ptr SSL, out_alert: ptr uint8 -): enum_ssl_verify_result_t {.cdecl.} = - let sslCtx = SSL_get_SSL_CTX(ssl) - - let quicCtx = cast[QuicContext](SSL_CTX_get_ex_data(sslCtx, SSL_CTX_ID)) - if quicCtx.isNil: - raiseAssert "could not obtain context" - - let derCertificates = getFullCertChain(ssl) - - let serverName = SSL_get_servername(ssl, TLSEXT_NAMETYPE_host_name) - doAssert quicCtx.tlsConfig.certVerifier.isSome, "no custom validator set" - if quicCtx.tlsConfig.certVerifier.get().verify($serverName, derCertificates): - return ssl_verify_ok - else: - out_alert[] = SSL_AD_CERTIFICATE_UNKNOWN +): enum_ssl_verify_result_t {.cdecl, raises: [].} = + try: + let sslCtx = SSL_get_SSL_CTX(ssl) + if sslCtx.isNil: + if not out_alert.isNil: + out_alert[] = SSL_AD_INTERNAL_ERROR + return ssl_verify_invalid + + let quicCtx = cast[QuicContext](SSL_CTX_get_ex_data(sslCtx, SSL_CTX_ID)) + if quicCtx.isNil or quicCtx.tlsConfig.certVerifier.isNone: + if not out_alert.isNil: + out_alert[] = SSL_AD_INTERNAL_ERROR + return ssl_verify_invalid + + let derCertificates = getFullCertChain(ssl) + let serverName = SSL_get_servername(ssl, TLSEXT_NAMETYPE_host_name) + if quicCtx.tlsConfig.certVerifier.get().verify($serverName, derCertificates): + return ssl_verify_ok + + if not out_alert.isNil: + out_alert[] = SSL_AD_CERTIFICATE_UNKNOWN + return ssl_verify_invalid + except Exception as exc: + warn "certificate verifier raised", errorMsg = exc.msg + if not out_alert.isNil: + out_alert[] = SSL_AD_CERTIFICATE_UNKNOWN return ssl_verify_invalid proc setupSSLContext*(quicCtx: QuicContext) = diff --git a/tests/test_verifier.nim b/tests/test_verifier.nim index 0b93426..8f8d33f 100644 --- a/tests/test_verifier.nim +++ b/tests/test_verifier.nim @@ -33,6 +33,13 @@ proc rejectingCertificateCb( discard derCertificates false +proc raisingCertificateCb( + serverName: string, derCertificates: seq[seq[byte]] +): bool {.gcsafe.} = + discard serverName + discard derCertificates + raise newException(ValueError, "verifier failed") + proc makeRejectingCertificateCb( recorder: RejectVerifierRecorder ): certificateVerifierCB = @@ -92,6 +99,18 @@ suite "certificate verifier": expect DialError: discard await client.dial(listener.localAddress()) + asyncTest "raising client verifier rejects handshake": + let client = + makeClientWithVerifier(CustomCertificateVerifier.init(raisingCertificateCb)) + let server = + makeServerWithVerifier(CustomCertificateVerifier.init(acceptingCertificateCb)) + let listener = server.listen(initTAddress("127.0.0.1:0")) + defer: + await allFutures(client.stop(), listener.stop()) + + expect DialError: + discard await client.dial(listener.localAddress()) + asyncTest "alpn mismatch rejects handshake": let client = makeClientWithVerifier( CustomCertificateVerifier.init(acceptingCertificateCb), singleAlpn("client-proto")