Skip to content
Open
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
2 changes: 2 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -9,3 +9,5 @@
/output

__pycache__

*.egg-info
42 changes: 42 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
[project]
name = "aws-cloudfront-authorizer"
version = "1.0.0"
requires-python = ">=3.12"
dependencies = [
"invoke>=2.2.1",
"pytest",
"troposphere>=4.10.1",
]

[tool.uv]
package = true

[tool.uv.sources]
central-helpers = { git = "ssh://git@bitbucket.org/vrt-prod/aws-cloudformation-helpers.git" }
custom-resources = { git = "https://github.com/vrtdev/custom-resources.git" }

[tool.setuptools]
packages = [
"templates",
]

[tool.ruff]
line-length = 140
indent-width = 4

[tool.ruff.lint]
extend-select = [
"E",
"W",
"A",
"COM",
"TID",
"B",
"SIM",
"UP",
]

[[tool.uv.index]]
name = "nexus"
url = "https://nexus.core.a51.be/repository/pypi/simple"
default = true
49 changes: 28 additions & 21 deletions src/authenticate.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,17 +6,25 @@
import jwt
import requests
import requests.auth

from cognito_utils import validate_cognito_id_token
from utils import bad_request, internal_server_error, get_refresh_token_jwt_secret, get_state_jwt_secret, \
generate_cookie, get_config
from aws_lambda_powertools import Logger
from aws_lambda_powertools.utilities.typing import LambdaContext

from cognito_utils import validate_cognito_id_token
from utils import (
bad_request,
generate_cookie,
get_config,
get_refresh_token_jwt_secret,
get_state_jwt_secret,
internal_server_error,
)

logger = Logger()

class InternalServerError(Exception): pass
class BadRequest(Exception): pass
class InternalServerError(Exception):
pass
class BadRequest(Exception):
pass


def exchange_cognito_code(event: dict, cognito_code: str) -> dict:
Expand Down Expand Up @@ -46,11 +54,11 @@ def exchange_cognito_code(event: dict, cognito_code: str) -> dict:
token_response = requests.post(
endpointurl,
data=post_data,
auth=requests.auth.HTTPBasicAuth(client_id, client_secret)
auth=requests.auth.HTTPBasicAuth(client_id, client_secret),
)
except requests.exceptions.ConnectionError as e:
logger.exception({"message": "Connection error to Cognito", "exception": e})

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is het de bedoeling om dit stukje logging te verliezen? Ik zie dezelfde change op nog plaatsen

raise InternalServerError()
except requests.exceptions.ConnectionError:
logger.exception({"message": "Connection error to Cognito"})
raise InternalServerError() from None

if token_response.status_code != 200:
try:
Expand All @@ -67,8 +75,7 @@ def exchange_cognito_code(event: dict, cognito_code: str) -> dict:
logger.exception({
"message": "Uncaught error",
"cognito_reply": token_response.text,
"exception": e,
"backtrace": traceback.format_exc()
"backtrace": traceback.format_exc(),
})
raise InternalServerError() from e

Expand All @@ -82,12 +89,12 @@ def exchange_cognito_code(event: dict, cognito_code: str) -> dict:
user_pool_id=os.environ['COGNITO_USER_POOL_ID'],
client_id=client_id,
)
except requests.exceptions.RequestException as e:
logger.exception({"message": "Connection error to Cognito", "exception": e})
raise InternalServerError()
except jwt.InvalidTokenError as e:
logger.exception({"message": "id_token invalid", "exception": e})
raise InternalServerError()
except requests.exceptions.RequestException:
logger.exception({"message": "Connection error to Cognito"})
raise InternalServerError() from None
except jwt.InvalidTokenError:
logger.exception({"message": "id_token invalid"})
raise InternalServerError() from None

logger.info("Cognito ID token is valid")

Expand Down Expand Up @@ -154,8 +161,8 @@ def handler(event, context: LambdaContext) -> dict:
f"redirect_uri={urllib.parse.quote_plus(state['redirect_uri'])}"
else:
raise ValueError(f"Invalid action `{state['action']}`")
except (KeyError, ValueError) as e:
logger.exception({"message": "state is invalid", "exception": e})
except (KeyError, ValueError):
logger.exception({"message": "state is invalid"})
return internal_server_error()

return {
Expand All @@ -166,7 +173,7 @@ def handler(event, context: LambdaContext) -> dict:
'Set-Cookie': generate_cookie(
get_config().cookie_name_refresh_token,
raw_refresh_token,
max_age=int(cognito_token['exp'] - now)
max_age=int(cognito_token['exp'] - now),
),
},
'body': 'Redirecting...',
Expand Down
30 changes: 20 additions & 10 deletions src/authorize.py
Original file line number Diff line number Diff line change
@@ -1,15 +1,25 @@
import time
from urllib.parse import urlsplit, urlunsplit, urlencode
from urllib.parse import urlencode, urlsplit, urlunsplit

import jwt
from aws_lambda_powertools import Logger
from aws_lambda_powertools.utilities.typing import LambdaContext

from utils import (
BadRequest,
InternalServerError,
NotLoggedIn,
access_token_from_refresh_token,
bad_request,
get_config,
get_refresh_token,
get_state_jwt_secret,
internal_server_error,
is_allowed_domain,
redirect_to_cognito,
)

logger = Logger()

from utils import get_config, bad_request, get_access_token_jwt_secret, redirect_to_cognito, NotLoggedIn, BadRequest, \
InternalServerError, internal_server_error, get_refresh_token, get_state_jwt_secret, is_allowed_domain, \
access_token_from_refresh_token

@logger.inject_lambda_context
def handler(event, context: LambdaContext) -> dict:
Expand Down Expand Up @@ -45,15 +55,15 @@ def handler(event, context: LambdaContext) -> dict:
logger.error(f"{redirect_uri} is not an allowed domain")
return bad_request('', f"{redirect_uri} is not an allowed domain")

if 'domains' in refresh_token: # delegated token with domain restrictions

@nielslaukens nielslaukens Sep 10, 2026

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Ik weet dat RUFF nogal opinionated is, maar ik vind de twee if-statements duidelijker dan de gecombineerde. De comment vond ik ook beter staan na de if condition: # then we are in this case

if redirect_uri_comp.netloc not in refresh_token['domains']:
logger.error(f"{redirect_uri} is not an allowed domain for this refresh token")
return bad_request('', f"{redirect_uri} is not an allowed domain for this refresh token")
# delegated token with domain restrictions
if 'domains' in refresh_token and redirect_uri_comp.netloc not in refresh_token['domains']:
logger.error(f"{redirect_uri} is not an allowed domain for this refresh token")
return bad_request('', f"{redirect_uri} is not an allowed domain for this refresh token")

try:
access_token = access_token_from_refresh_token(
refresh_token,
redirect_uri_comp.netloc
redirect_uri_comp.netloc,
)
except BadRequest as e:
return bad_request('', e)
Expand Down
20 changes: 13 additions & 7 deletions src/batch_authorize.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,19 @@
import json

from utils import bad_request, NotLoggedIn, BadRequest, \
InternalServerError, internal_server_error, get_refresh_token, get_domains, \
access_token_from_refresh_token
from aws_lambda_powertools import Logger
from aws_lambda_powertools.utilities.typing import LambdaContext

from utils import (
BadRequest,
InternalServerError,
NotLoggedIn,
access_token_from_refresh_token,
bad_request,
get_domains,
get_refresh_token,
internal_server_error,
)

logger = Logger()

@logger.inject_lambda_context
Expand All @@ -25,10 +33,8 @@ def handler(event, context: LambdaContext) -> dict:
except InternalServerError as e:
return internal_server_error('', e)

if 'domains' in refresh_token: # delegated token with domain restrictions

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Dit vind het origineel hier ook duidelijker dan de ruff-versie

domains = refresh_token['domains']
else:
domains = get_domains()
# delegated token with domain restrictions
domains = refresh_token['domains'] if 'domains' in refresh_token else get_domains()

access_tokens = {}
try:
Expand Down
5 changes: 2 additions & 3 deletions src/cognito_utils.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import contextlib
import functools

import jwt
Expand All @@ -6,10 +7,8 @@
# Use pure python implementation for crypto
from jwt_rsa_algo import RsaAlgorithm

try:
with contextlib.suppress(ValueError): # Assume already registered
jwt.register_algorithm('RS256', RsaAlgorithm(RsaAlgorithm.SHA256))
except ValueError:
pass # Assume already registered


@functools.lru_cache(maxsize=1)
Expand Down
28 changes: 19 additions & 9 deletions src/delegate.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,13 +5,25 @@
import urllib.parse

import jwt

from utils import redirect_to_cognito, get_refresh_token, NotLoggedIn, BadRequest, \
bad_request, InternalServerError, internal_server_error, \
get_grant_jwt_secret, get_state_jwt_secret, get_config, is_allowed_domain, dynamodb_client, get_domains
from aws_lambda_powertools import Logger
from aws_lambda_powertools.utilities.typing import LambdaContext

from utils import (
BadRequest,
InternalServerError,
NotLoggedIn,
bad_request,
dynamodb_client,
get_config,
get_domains,
get_grant_jwt_secret,
get_refresh_token,
get_state_jwt_secret,
internal_server_error,
is_allowed_domain,
redirect_to_cognito,
)

logger = Logger()

@logger.inject_lambda_context
Expand Down Expand Up @@ -55,9 +67,8 @@ def handler(event, context: LambdaContext) -> dict:
for group_entry in page['Items']:
try:
groups[group_entry['group']['S']] = group_entry['domains']['SS']
except KeyError as e:
except KeyError:
logger.exception("Invalid group in DynamoDB: " + repr(group_entry))
pass

with open(os.path.join(os.path.dirname(__file__), 'delegate.html')) as f:
html = f.read()
Expand Down Expand Up @@ -100,9 +111,8 @@ def handler(event, context: LambdaContext) -> dict:
if not is_allowed_domain(domain):
return bad_request('', 'Unknown domain in request')

if 'domains' in refresh_token:
if not domains.issubset(refresh_token['domains']):
return bad_request('', 'domain requested outside refresh_token')
if 'domains' in refresh_token and not domains.issubset(refresh_token['domains']):
return bad_request('', 'domain requested outside refresh_token')

# Validate no commas in new subject to avoid future join ambiguity
if ',' in subject:
Expand Down
4 changes: 2 additions & 2 deletions src/generate_ci.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,11 +3,11 @@
import time

import jwt

from utils import get_access_token_jwt_secret, bad_request, is_allowed_domain
from aws_lambda_powertools import Logger
from aws_lambda_powertools.utilities.typing import LambdaContext

from utils import bad_request, get_access_token_jwt_secret, is_allowed_domain

logger = Logger()

@logger.inject_lambda_context
Expand Down
18 changes: 13 additions & 5 deletions src/index.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,21 @@
import contextlib
import json
import os
import time

import jwt

from utils import NotLoggedIn, BadRequest, InternalServerError, internal_server_error, cognito_url, \
get_state_jwt_secret, get_csrf_jwt_secret, get_raw_refresh_token, parse_raw_refresh_token
from utils import (
BadRequest,
InternalServerError,
NotLoggedIn,
cognito_url,
get_csrf_jwt_secret,
get_raw_refresh_token,
get_state_jwt_secret,
internal_server_error,
parse_raw_refresh_token,
)


def handler(event, context) -> dict:
Expand All @@ -24,10 +34,8 @@ def handler(event, context) -> dict:
azp = refresh_token['azp'] # Mandatory
sub = refresh_token.get('sub', []) # optional

try:
with contextlib.suppress(KeyError):
domains = refresh_token['domains']
except KeyError:
pass
except (NotLoggedIn, BadRequest):
pass
except InternalServerError as e:
Expand Down
9 changes: 8 additions & 1 deletion src/logout.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,14 @@
from aws_lambda_powertools import Logger
from aws_lambda_powertools.utilities.typing import LambdaContext

from utils import generate_cookie, get_config, bad_request, get_csrf_jwt_secret, get_raw_refresh_token, NotLoggedIn
from utils import (
NotLoggedIn,
bad_request,
generate_cookie,
get_config,
get_csrf_jwt_secret,
get_raw_refresh_token,
)

logger = Logger()

Expand Down
10 changes: 8 additions & 2 deletions src/use_grant.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,13 @@

import jwt

from utils import bad_request, get_grant_jwt_secret, generate_cookie, get_config, get_refresh_token_jwt_secret
from utils import (
bad_request,
generate_cookie,
get_config,
get_grant_jwt_secret,
get_refresh_token_jwt_secret,
)


def handler(event, context) -> dict:
Expand Down Expand Up @@ -33,7 +39,7 @@ def handler(event, context) -> dict:
raw_refresh_token = jwt.encode(
refresh_token,
get_refresh_token_jwt_secret(),
algorithm='HS256'
algorithm='HS256',
)

with open(os.path.join(os.path.dirname(__file__), 'use_grant.html')) as f:
Expand Down
Loading