@@ -86,28 +86,66 @@ def test_verify_token(_fetch_certs, decode):
8686
8787
8888@mock .patch ("google.oauth2.id_token._fetch_certs" , autospec = True )
89- @mock .patch ("jwt.PyJWKClient" , autospec = True )
89+ @mock .patch ("jwt.api_jwk.PyJWKSet" , autospec = True )
90+ @mock .patch ("jwt.get_unverified_header" , autospec = True )
9091@mock .patch ("jwt.decode" , autospec = True )
91- def test_verify_token_jwk (decode , py_jwk , _fetch_certs ):
92+ def test_verify_token_jwk (decode , get_unverified_header , py_jwk_set , _fetch_certs ):
9293 certs_url = "abc123"
9394 data = {"keys" : [{"alg" : "RS256" }]}
9495 _fetch_certs .return_value = data
96+ get_unverified_header .return_value = {"kid" : "mock-kid" }
97+
98+ mock_key = mock .MagicMock ()
99+ mock_key .key_id = "mock-kid"
100+ mock_key .public_key_use = "sig"
101+ mock_key .key = mock .sentinel .key
102+ mock_key .algorithm_name = "mock-alg"
103+ py_jwk_set .from_dict .return_value .keys = [mock_key ]
95104 result = id_token .verify_token (
96105 mock .sentinel .token , mock .sentinel .request , certs_url = certs_url
97106 )
98107 assert result == decode .return_value
99- py_jwk .assert_called_once_with (certs_url )
100- signing_key = py_jwk .return_value .get_signing_key_from_jwt
108+ py_jwk_set .from_dict .assert_called_once_with (data )
109+ get_unverified_header .assert_called_once_with (mock .sentinel .token )
110+
101111 _fetch_certs .assert_called_once_with (mock .sentinel .request , certs_url )
102- signing_key .assert_called_once_with (mock .sentinel .token )
103112 decode .assert_called_once_with (
104113 mock .sentinel .token ,
105- signing_key . return_value .key ,
106- algorithms = [signing_key . return_value . algorithm_name ],
114+ mock . sentinel .key ,
115+ algorithms = ["mock-alg" ],
107116 audience = None ,
108117 )
109118
110119
120+ @mock .patch ("google.oauth2.id_token._fetch_certs" , autospec = True )
121+ @mock .patch ("jwt.api_jwk.PyJWKSet" , autospec = True )
122+ @mock .patch ("jwt.get_unverified_header" , autospec = True )
123+ @mock .patch ("jwt.decode" , autospec = True )
124+ def test_verify_token_jwk_missing_kid (
125+ decode , get_unverified_header , py_jwk_set , _fetch_certs
126+ ):
127+ from jwt .exceptions import PyJWKClientError
128+
129+ certs_url = "abc123"
130+ data = {"keys" : [{"alg" : "RS256" }]}
131+ _fetch_certs .return_value = data
132+ get_unverified_header .return_value = {"kid" : "mock-kid" }
133+
134+ mock_key = mock .MagicMock ()
135+ mock_key .key_id = "different-kid"
136+ mock_key .public_key_use = "sig"
137+ mock_key .key = mock .sentinel .key
138+ mock_key .algorithm_name = "mock-alg"
139+ py_jwk_set .from_dict .return_value .keys = [mock_key ]
140+
141+ with pytest .raises (
142+ PyJWKClientError , match = 'Unable to find a signing key that matches: "mock-kid"'
143+ ):
144+ id_token .verify_token (
145+ mock .sentinel .token , mock .sentinel .request , certs_url = certs_url
146+ )
147+
148+
111149@mock .patch ("google.auth.jwt.decode" , autospec = True )
112150@mock .patch ("google.oauth2.id_token._fetch_certs" , autospec = True )
113151def test_verify_token_args (_fetch_certs , decode ):
0 commit comments