Skip to content
Merged
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
211 changes: 131 additions & 80 deletions src/mock_vws/_model_target_web_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -482,94 +482,19 @@ def encode_part(value: dict[str, JSONValue]) -> str:


@beartype
def oauth2_token( # noqa: PLR0911 # pylint: disable=too-many-return-statements
def _oauth2_access_token_response(
*,
request: RequestData,
credential_store: ModelTargetDatasetStore,
form: dict[str, list[str]],
auth_header: str | None,
credential_scopes: frozenset[str],
) -> _ResponseType:
"""Return a fake OAuth2 access token."""
content_length_error = _content_length_error(request=request)
if content_length_error is not None:
return content_length_error

auth_header = _get_header(request=request, name="Authorization")
# A form body which is not valid UTF-8 is decoded leniently rather than
# raising, so that a body which cannot be decoded is treated as one which
# does not name a grant type.
form = parse_qs(
qs=request.body.decode(encoding="utf-8", errors="replace"),
)
grant_type = form.get("grant_type", ["client_credentials"])[0]
if grant_type not in {"client_credentials", "password"}:
return _oauth2_error_response(
status_code=HTTPStatus.BAD_REQUEST,
body={"error": "unsupported_grant_type"},
)

dynamic_credential: OAuth2ClientCredential | None = None
if grant_type == "client_credentials":
basic_credentials = _basic_auth_credentials(auth_header=auth_header)
if basic_credentials is None:
return _oauth2_error_response(
status_code=HTTPStatus.UNAUTHORIZED,
body={
"error": "invalid_request",
"error_description": (
"Missing or invalid authorization header"
),
},
)

dynamic_credential = credential_store.oauth2_client_credentials.get(
basic_credentials[0],
)
fixed_credential_matches = basic_credentials == (
_MOCK_MODEL_TARGET_CLIENT_ID,
_MOCK_MODEL_TARGET_CLIENT_SECRET,
)
dynamic_credential_matches = (
dynamic_credential is not None
and dynamic_credential.client_secret == basic_credentials[1]
)
if not fixed_credential_matches and not dynamic_credential_matches:
return _oauth2_error_response(
status_code=HTTPStatus.UNAUTHORIZED,
body={"error": "invalid_client"},
)
else:
username = form.get("username", [""])[0]
password = form.get("password", [""])[0]
if len(username) == 0 or len(password) == 0:
return _oauth2_error_response(
status_code=HTTPStatus.BAD_REQUEST,
body={
"error": "invalid_request",
"error_description": "Missing username and/or password",
},
)
if (username, password) != (
_MOCK_MODEL_TARGET_USERNAME,
_MOCK_MODEL_TARGET_PASSWORD,
):
return _oauth2_error_response(
status_code=HTTPStatus.UNAUTHORIZED,
body={
"error": "invalid_grant",
"error_description": "Invalid username and/or password",
},
)

"""Return an access token limited to ``credential_scopes``."""
auth_text = auth_header if auth_header is not None else ""
token_source = (
request.body if len(request.body) > 0 else auth_text.encode()
)
requested_scope = form.get("scope", [""])[0]
if grant_type == "client_credentials" and dynamic_credential is not None:
credential_scopes = frozenset(dynamic_credential.scopes)
else:
credential_scopes = _MODEL_TARGET_SCOPES | {
_CLIENT_CREDENTIALS_SCOPE,
}
requested_scopes = frozenset(requested_scope.split())
scopes = (
requested_scopes if len(requested_scopes) > 0 else credential_scopes
Expand All @@ -592,6 +517,132 @@ def oauth2_token( # noqa: PLR0911 # pylint: disable=too-many-return-statements
)


@beartype
def _oauth2_client_credentials_token(
*,
request: RequestData,
form: dict[str, list[str]],
auth_header: str | None,
credential_store: ModelTargetDatasetStore,
) -> _ResponseType:
"""Validate a client-credentials grant and return its token
response.
"""
basic_credentials = _basic_auth_credentials(auth_header=auth_header)
if basic_credentials is None:
return _oauth2_error_response(
status_code=HTTPStatus.UNAUTHORIZED,
body={
"error": "invalid_request",
"error_description": "Missing or invalid authorization header",
},
)

dynamic_credential = credential_store.oauth2_client_credentials.get(
basic_credentials[0],
)
fixed_credential_matches = basic_credentials == (
_MOCK_MODEL_TARGET_CLIENT_ID,
_MOCK_MODEL_TARGET_CLIENT_SECRET,
)
dynamic_credential_matches = (
dynamic_credential is not None
and dynamic_credential.client_secret == basic_credentials[1]
)
if not fixed_credential_matches and not dynamic_credential_matches:
return _oauth2_error_response(
status_code=HTTPStatus.UNAUTHORIZED,
body={"error": "invalid_client"},
)
credential_scopes = (
frozenset(dynamic_credential.scopes)
if dynamic_credential is not None
else _MODEL_TARGET_SCOPES | {_CLIENT_CREDENTIALS_SCOPE}
)
return _oauth2_access_token_response(
request=request,
form=form,
auth_header=auth_header,
credential_scopes=credential_scopes,
)


@beartype
def _oauth2_password_token(
*,
request: RequestData,
form: dict[str, list[str]],
auth_header: str | None,
) -> _ResponseType:
"""Validate a password grant and return its token response."""
username = form.get("username", [""])[0]
password = form.get("password", [""])[0]
if len(username) == 0 or len(password) == 0:
return _oauth2_error_response(
status_code=HTTPStatus.BAD_REQUEST,
body={
"error": "invalid_request",
"error_description": "Missing username and/or password",
},
)
if (username, password) != (
_MOCK_MODEL_TARGET_USERNAME,
_MOCK_MODEL_TARGET_PASSWORD,
):
return _oauth2_error_response(
status_code=HTTPStatus.UNAUTHORIZED,
body={
"error": "invalid_grant",
"error_description": "Invalid username and/or password",
},
)
return _oauth2_access_token_response(
request=request,
form=form,
auth_header=auth_header,
credential_scopes=(_MODEL_TARGET_SCOPES | {_CLIENT_CREDENTIALS_SCOPE}),
)


@beartype
def oauth2_token(
*,
request: RequestData,
credential_store: ModelTargetDatasetStore,
) -> _ResponseType:
"""Return a fake OAuth2 access token."""
content_length_error = _content_length_error(request=request)
if content_length_error is not None:
return content_length_error

auth_header = _get_header(request=request, name="Authorization")
# A form body which is not valid UTF-8 is decoded leniently rather than
# raising, so that a body which cannot be decoded is treated as one which
# does not name a grant type.
form = parse_qs(
qs=request.body.decode(encoding="utf-8", errors="replace"),
)
grant_type = form.get("grant_type", ["client_credentials"])[0]
if grant_type not in {"client_credentials", "password"}:
return _oauth2_error_response(
status_code=HTTPStatus.BAD_REQUEST,
body={"error": "unsupported_grant_type"},
)

if grant_type == "client_credentials":
return _oauth2_client_credentials_token(
request=request,
form=form,
auth_header=auth_header,
credential_store=credential_store,
)
return _oauth2_password_token(
request=request,
form=form,
auth_header=auth_header,
)


@beartype
def _require_client_credentials_scope(
request: RequestData,
Expand Down
Loading