Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 8 additions & 2 deletions synapse_token_authenticator/claims_validator.py
Original file line number Diff line number Diff line change
Expand Up @@ -163,12 +163,18 @@ def parse_validator(d: dict | list) -> Validator:
val_type = d.pop("type")
validator = VALIDATORS.get(val_type)
if validator:
return validator(**d)
try:
return validator(**d)
except TypeError as e:
raise InvalidClaimsValidatorError(f"Invalid validator arguments: {e}")
raise InvalidClaimsValidatorError(f"Unknown validator type {val_type}")
if isinstance(d, list):
val_type = d.pop(0)
validator = VALIDATORS.get(val_type)
if validator:
return validator(*d)
try:
return validator(*d)
except TypeError as e:
raise InvalidClaimsValidatorError(f"Invalid validator arguments: {e}")
raise InvalidClaimsValidatorError(f"Unknown validator type {val_type}")
raise InvalidClaimsValidatorError("Validator parsing failed, expected list or dict")
58 changes: 34 additions & 24 deletions synapse_token_authenticator/config/epa.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,8 @@ def parse_enc_jwk(cls, value: Any) -> JWK | None:
return None
if isinstance(value, JWK):
return value
if isinstance(value, str):
return JWK.from_json(value)
if isinstance(value, dict):
return JWK(**value)
return None
Expand All @@ -74,11 +76,15 @@ def parse_jwk_set(cls, value: Any) -> JWKSet | JWK | None:
if isinstance(value, (JWKSet, JWK)):
return value
if isinstance(value, str):
return JWKSet.from_json(value)
if isinstance(value, dict) and "keys" in value:
return JWKSet.from_json(json.dumps(value))
if json.loads(value).get("keys"):
return JWKSet.from_json(value)
else:
return JWK.from_json(value)
if isinstance(value, dict):
return JWK(**value)
if "keys" in value:
return JWKSet.from_json(json.dumps(value))
else:
return JWK(**value)
return None

@model_validator(mode="after")
Expand All @@ -87,15 +93,17 @@ def decide_enc_jwk(self) -> Self:
self.enc_jwk is not None,
self.enc_jwk_file is not None,
]
if sum(sources) != 1:
raise ValueError("Exactly one of enc_jwk or enc_jwk_file must be set")
if self.enc_jwk:
return self
elif self.enc_jwk_file:
with open(self.enc_jwk_file, "rb") as f:
self.enc_jwk = JWK.from_pem(f.read())
if sum(sources) == 1:
if self.enc_jwk:
return self
raise ValueError("No encryption JWK")
elif self.enc_jwk_file:
try:
with open(self.enc_jwk_file, "rb") as f:
self.enc_jwk = JWK.from_pem(f.read())
return self
except FileNotFoundError:
raise ValueError(f"enc_jwk file '{self.enc_jwk_file}' not found")
raise ValueError("Exactly one of enc_jwk or enc_jwk_file must be set")

@model_validator(mode="after")
def decide_jwk_set(self) -> Self:
Expand All @@ -104,16 +112,18 @@ def decide_jwk_set(self) -> Self:
self.jwk_file is not None,
self.jwks_endpoint is not None,
]
if sum(sources) != 1:
raise ValueError(
"Exactly one of jwk_set, jwk_file, or jwks_endpoint must be set"
)
if self.jwk_set:
return self
elif self.jwk_file:
with open(self.jwk_file, "rb") as f:
self.jwk_set = JWK.from_pem(f.read())
if sum(sources) == 1:
if self.jwk_set:
return self
elif self.jwk_file:
try:
with open(self.jwk_file, "rb") as f:
self.jwk_set = JWK.from_pem(f.read())
return self
except FileNotFoundError:
raise ValueError(f"jwk_file '{self.jwk_file}' not found")
elif self.jwks_endpoint:
return self
elif self.jwks_endpoint:
return self
raise ValueError("No JWK set")
raise ValueError(
"Exactly one of jwk_set, jwk_file, or jwks_endpoint must be set"
)
38 changes: 22 additions & 16 deletions synapse_token_authenticator/config/oauth.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,11 +56,15 @@ def parse_jwk_set(cls, value: Any) -> JWKSet | JWK | None:
if isinstance(value, (JWKSet, JWK)):
return value
if isinstance(value, str):
return JWKSet.from_json(value)
if isinstance(value, dict) and "keys" in value:
return JWKSet.from_json(json.dumps(value))
if json.loads(value).get("keys"):
return JWKSet.from_json(value)
else:
return JWK.from_json(value)
if isinstance(value, dict):
return JWK(**value)
if "keys" in value:
return JWKSet.from_json(json.dumps(value))
else:
return JWK(**value)
return None

@model_validator(mode="after")
Expand All @@ -70,19 +74,21 @@ def decide_jwk_set(self) -> Self:
self.jwk_file is not None,
self.jwks_endpoint is not None,
]
if sum(sources) != 1:
raise ValueError(
"Exactly one of jwk_set, jwk_file, or jwks_endpoint must be set"
)
if self.jwk_set:
return self
elif self.jwk_file:
with open(self.jwk_file, "rb") as f:
self.jwk_set = JWK.from_pem(f.read())
if sum(sources) == 1:
if self.jwk_set:
return self
elif self.jwk_file:
try:
with open(self.jwk_file, "rb") as f:
self.jwk_set = JWK.from_pem(f.read())
return self
except FileNotFoundError:
raise ValueError(f"jwk_file '{self.jwk_file}' not found")
elif self.jwks_endpoint:
return self
elif self.jwks_endpoint:
return self
raise ValueError("No JWK set")
raise ValueError(
"Exactly one of jwk_set, jwk_file, or jwks_endpoint must be set"
)


@dataclass(config=ConfigDict(arbitrary_types_allowed=True, extra="ignore"))
Expand Down
4 changes: 4 additions & 0 deletions synapse_token_authenticator/token_authenticator.py
Original file line number Diff line number Diff line change
Expand Up @@ -101,6 +101,10 @@ def __init__(self, config: TokenAuthenticatorConfig, module_api: ModuleApi):

# Registers the encryption public keys
keys = JWKSet()

# enc_jwk and enc_jwk_file are both optional fields but either one must set.
# If enc_jwk is empty, the model resolves it from enc_jwk_file. So enc_jwk
# cannot be empty. This assert is to resolve mypy error.
assert self.config.epa.enc_jwk is not None
keys.add(self.config.epa.enc_jwk)
self.api.register_web_resource(
Expand Down
4 changes: 3 additions & 1 deletion tests/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -178,10 +178,12 @@ def get_jwt_token(
claims=None,
id_="123456",
extra_headers=None,
key=None,
) -> str:
if extra_headers is None:
extra_headers = {}
key = get_jwk(secret, id_)
if key is None:
key = get_jwk(secret, id_)
if claims is None:
claims = {}
claims["sub"] = username
Expand Down
Loading
Loading