44This handles validating messages sent by the tool and generating
55access token with LTI scopes.
66"""
7- import codecs
87import copy
9- import time
108import json
9+ import math
10+ import time
11+ import sys
1112
13+ import jwt
1214from Cryptodome .PublicKey import RSA
13- from jwkest import BadSignature , BadSyntax , WrongNumberOfParts , jwk
14- from jwkest .jwk import RSAKey , load_jwks_from_url
15- from jwkest .jws import JWS , NoSuitableSigningKeys
16- from jwkest .jwt import JWT
15+ from jwt .api_jwk import PyJWK
1716
1817from . import exceptions
1918
@@ -47,14 +46,9 @@ def __init__(self, public_key=None, keyset_url=None):
4746 # Import from public key
4847 if public_key :
4948 try :
50- new_key = RSAKey (use = 'sig' )
51-
52- # Unescape key before importing it
53- raw_key = codecs .decode (public_key , 'unicode_escape' )
54-
5549 # Import Key and save to internal state
56- new_key . load_key ( RSA . import_key ( raw_key ) )
57- self .public_key = new_key
50+ algo_obj = jwt . get_algorithm_by_name ( 'RS256' )
51+ self .public_key = PyJWK . from_json ( algo_obj . to_jwk ( public_key ))
5852 except ValueError as err :
5953 raise exceptions .InvalidRsaKey () from err
6054
@@ -69,7 +63,7 @@ def _get_keyset(self, kid=None):
6963
7064 if self .keyset_url :
7165 try :
72- keys = load_jwks_from_url (self .keyset_url )
66+ keys = jwt . PyJWKClient (self .keyset_url ). get_jwk_set ( )
7367 except Exception as err :
7468 # Broad Exception is required here because jwkest raises
7569 # an Exception object explicitly.
@@ -78,13 +72,13 @@ def _get_keyset(self, kid=None):
7872 raise exceptions .NoSuitableKeys () from err
7973 keyset .extend (keys )
8074
81- if self .public_key and kid :
82- # Fill in key id of stored key.
83- # This is needed because if the JWS is signed with a
84- # key with a kid, pyjwkest doesn't match them with
85- # keys without kid ( kid=None) and fails verification
86- self . public_key . kid = kid
87-
75+ if self .public_key :
76+ if kid :
77+ # Fill in key id of stored key.
78+ # This is needed because if the JWS is signed with a
79+ # key with a kid, pyjwkest doesn't match them with
80+ # keys without kid (kid=None) and fails verification
81+ self . public_key . kid = kid
8882 # Add to keyset
8983 keyset .append (self .public_key )
9084
@@ -100,32 +94,24 @@ def validate_and_decode(self, token):
10094 iss, sub, exp, aud and jti claims.
10195 """
10296 try :
103- # Get KID from JWT header
104- jwt = JWT ().unpack (token )
105-
106- # Verify message signature
107- message = JWS ().verify_compact (
108- token ,
109- keys = self ._get_keyset (
110- jwt .headers .get ('kid' )
111- )
112- )
113-
114- # If message is valid, check expiration from JWT
115- if 'exp' in message and message ['exp' ] < time .time ():
116- raise exceptions .TokenSignatureExpired ()
117-
118- # TODO: Validate other JWT claims
119-
120- # Else returns decoded message
121- return message
122-
123- except NoSuitableSigningKeys as err :
124- raise exceptions .NoSuitableKeys () from err
125- except (BadSyntax , WrongNumberOfParts ) as err :
126- raise exceptions .MalformedJwtToken () from err
127- except BadSignature as err :
128- raise exceptions .BadJwtSignature () from err
97+ key_set = self ._get_keyset ()
98+ if not key_set :
99+ raise exceptions .NoSuitableKeys ()
100+ for i in range (len (key_set )):
101+ try :
102+ message = jwt .decode (
103+ token ,
104+ key = key_set [i ],
105+ algorithms = ['RS256' , 'RS512' ,],
106+ options = {'verify_signature' : True }
107+ )
108+ return message
109+ except Exception :
110+ if i == len (key_set ) - 1 :
111+ raise
112+ except Exception as token_error :
113+ exc_info = sys .exc_info ()
114+ raise jwt .InvalidTokenError (exc_info [2 ]) from token_error
129115
130116
131117class PlatformKeyHandler :
@@ -144,14 +130,8 @@ def __init__(self, key_pem, kid=None):
144130 if key_pem :
145131 # Import JWK from RSA key
146132 try :
147- self .key = RSAKey (
148- # Using the same key ID as client id
149- # This way we can easily serve multiple public
150- # keys on teh same endpoint and keep all
151- # LTI 1.3 blocks working
152- kid = kid ,
153- key = RSA .import_key (key_pem )
154- )
133+ algo = jwt .get_algorithm_by_name ('RS256' )
134+ self .key = algo .prepare_key (key_pem )
155135 except ValueError as err :
156136 raise exceptions .InvalidRsaKey () from err
157137
@@ -167,28 +147,26 @@ def encode_and_sign(self, message, expiration=None):
167147 # Set iat and exp if expiration is set
168148 if expiration :
169149 _message .update ({
170- "iat" : int (round (time .time ())),
171- "exp" : int (round (time .time ()) + expiration ),
150+ "iat" : int (math . floor (time .time ())),
151+ "exp" : int (math . floor (time .time ()) + expiration ),
172152 })
173153
174154 # The class instance that sets up the signing operation
175155 # An RS 256 key is required for LTI 1.3
176- _jws = JWS (_message , alg = "RS256" , cty = "JWT" )
177-
178- # Encode and sign LTI message
179- return _jws .sign_compact ([self .key ])
156+ return jwt .encode (_message , self .key , algorithm = "RS256" )
180157
181158 def get_public_jwk (self ):
182159 """
183160 Export Public JWK
184161 """
185- public_keys = jwk . KEYS ()
162+ jwk = { "keys" : []}
186163
187164 # Only append to keyset if a key exists
188165 if self .key :
189- public_keys .append (self .key )
190-
191- return json .loads (public_keys .dump_jwks ())
166+ algo_obj = jwt .get_algorithm_by_name ('RS256' )
167+ public_key = algo_obj .prepare_key (self .key ).public_key ()
168+ jwk ['keys' ].append (json .loads (algo_obj .to_jwk (public_key )))
169+ return jwk
192170
193171 def validate_and_decode (self , token , iss = None , aud = None ):
194172 """
@@ -197,29 +175,22 @@ def validate_and_decode(self, token, iss=None, aud=None):
197175 Validates a token sent by the tool using the platform's RSA Key.
198176 Optionally validate iss and aud claims if provided.
199177 """
178+ if not self .key :
179+ raise exceptions .RsaKeyNotSet ()
200180 try :
201- # Verify message signature
202- message = JWS ().verify_compact (token , keys = [self .key ])
203-
204- # If message is valid, check expiration from JWT
205- if 'exp' in message and message ['exp' ] < time .time ():
206- raise exceptions .TokenSignatureExpired ()
207-
208- # Validate issuer claim (if present)
209- if iss :
210- if 'iss' not in message or message ['iss' ] != iss :
211- raise exceptions .InvalidClaimValue ('The required iss claim is either missing or does '
212- 'not match the expected iss value.' )
213-
214- # Validate audience claim (if present)
215- if aud :
216- if 'aud' not in message or aud not in message ['aud' ]:
217- raise exceptions .InvalidClaimValue ('The required aud claim is missing.' )
218-
219- # Else return token contents
181+ message = jwt .decode (
182+ token ,
183+ key = self .key .public_key (),
184+ audience = aud ,
185+ issuer = iss ,
186+ algorithms = ['RS256' , 'RS512' ],
187+ options = {
188+ 'verify_signature' : True ,
189+ 'verify_aud' : True if aud else False
190+ }
191+ )
220192 return message
221193
222- except NoSuitableSigningKeys as err :
223- raise exceptions .NoSuitableKeys () from err
224- except BadSyntax as err :
225- raise exceptions .MalformedJwtToken () from err
194+ except Exception as token_error :
195+ exc_info = sys .exc_info ()
196+ raise jwt .InvalidTokenError (exc_info [2 ]) from token_error
0 commit comments