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
8 changes: 2 additions & 6 deletions src/mock_vws/_flask_server/vwq.py
Original file line number Diff line number Diff line change
Expand Up @@ -137,7 +137,7 @@ def query() -> Response:

databases = get_all_cloud_databases()
request_body = request.stream.read()
run_query_validators(
validated_query = run_query_validators(
request_headers=dict(request.headers),
request_body=request_body,
request_method=request.method,
Expand All @@ -147,11 +147,7 @@ def query() -> Response:
date = email.utils.formatdate(timeval=None, localtime=False, usegmt=True)

response_text = get_query_match_response_text(
request_headers=dict(request.headers),
request_body=request_body,
request_method=request.method,
request_path=request.path,
databases=databases,
validated_query=validated_query,
query_match_checker=query_match_checker,
)

Expand Down
44 changes: 9 additions & 35 deletions src/mock_vws/_query_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,65 +2,39 @@

import base64
import uuid
from collections.abc import Iterable, Mapping
from typing import Any

from beartype import beartype

from mock_vws._base64_decoding import decode_base64
from mock_vws._constants import ResultCodes, TargetStatuses
from mock_vws._database_matchers import get_database_matching_client_keys
from mock_vws._matching import matching_targets
from mock_vws._mock_common import json_dump
from mock_vws._query_validators.multipart import parse_multipart
from mock_vws.database import CloudDatabase
from mock_vws._query_validators import ValidatedQuery
from mock_vws.image_matchers import ImageMatcher


@beartype
def get_query_match_response_text(
*,
request_headers: Mapping[str, str],
request_body: bytes,
request_method: str,
request_path: str,
databases: Iterable[CloudDatabase],
validated_query: ValidatedQuery,
query_match_checker: ImageMatcher,
) -> str:
"""
Args:
request_path: The path of the request.
request_headers: The headers sent with the request.
request_body: The body of the request.
request_method: The HTTP method of the request.
databases: All Vuforia databases.
validated_query: The database and the parsed body which the query
validators resolved the request to.
query_match_checker: A callable which takes two image values and
returns a match score, or ``None`` if they do not match.

Returns:
The response text for a query endpoint request.
"""
fields, files = parse_multipart(
request_headers=request_headers,
request_body=request_body,
)

max_num_results = fields.get(key="max_num_results", default="1")
include_target_data = fields.get(
key="include_target_data",
default="top",
).lower()

image_part = files["image"]
image_value = image_part.stream.read()

database = get_database_matching_client_keys(
request_headers=request_headers,
request_body=request_body,
request_method=request_method,
request_path=request_path,
databases=databases,
)
fields = validated_query.form.fields
max_num_results = fields.get("max_num_results", "1")
include_target_data = fields.get("include_target_data", "top").lower()
image_value = validated_query.form.files["image"]
database = validated_query.database

matches_best_first = matching_targets(
matcher=query_match_checker,
Expand Down
87 changes: 49 additions & 38 deletions src/mock_vws/_query_validators/__init__.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,14 @@
"""Input validators to use in the mock query API."""

from collections.abc import Iterable, Mapping
from dataclasses import dataclass

from beartype import beartype

from mock_vws._query_validators.multipart import (
MultipartForm,
parse_multipart,
)
from mock_vws.database import CloudDatabase

from .accept_header_validators import validate_accept_header
Expand Down Expand Up @@ -38,6 +43,24 @@
from .project_state_validators import validate_project_state


@beartype
@dataclass(frozen=True, kw_only=True)
class ValidatedQuery:
"""What the validators learn about a query request which passes them.

Args:
database: The database which the request's client keys belong to.
form: The parsed body of the request.

Attributes:
database: The database which the request's client keys belong to.
form: The parsed body of the request.
"""

database: CloudDatabase
form: MultipartForm


@beartype
def run_query_validators(
*,
Expand All @@ -46,15 +69,28 @@ def run_query_validators(
request_body: bytes,
request_method: str,
databases: Iterable[CloudDatabase],
) -> None:
) -> ValidatedQuery:
"""Run all validators.

Vuforia reports one problem with a request even when the request has
more than one. Which problem it reports is decided by the order of the
validators here, so that order is the mock's record of Vuforia's error
precedence, verified against the real service.

The body is parsed once, after the ``Content-Type`` header which names
its boundary has been validated, and the parsed form is shared by every
validator which reads the body.

Args:
request_path: The path of the request.
request_headers: The headers sent with the request.
request_body: The body of the request.
request_method: The HTTP method of the request.
databases: All Vuforia databases.

Returns:
The database which the request's client keys belong to, and the
parsed body of the request.
"""
validate_content_length_header_is_int(request_headers=request_headers)
validate_content_length_header_not_too_large(
Expand All @@ -72,20 +108,14 @@ def run_query_validators(
request_headers=request_headers,
databases=databases,
)
validate_authorization(
request_headers=request_headers,
request_body=request_body,
request_method=request_method,
request_path=request_path,
databases=databases,
)
validate_project_state(
database = validate_authorization(
request_headers=request_headers,
request_body=request_body,
request_method=request_method,
request_path=request_path,
databases=databases,
)
validate_project_state(database=database)
validate_accept_header(request_headers=request_headers)
validate_date_header_given(request_headers=request_headers)
validate_date_format(request_headers=request_headers)
Expand All @@ -94,35 +124,16 @@ def run_query_validators(
request_headers=request_headers,
request_body=request_body,
)
validate_extra_fields(
request_headers=request_headers,
request_body=request_body,
)
validate_image_field_given(
request_headers=request_headers,
request_body=request_body,
)
validate_image_is_image(
request_headers=request_headers,
request_body=request_body,
)
validate_image_format(
request_headers=request_headers,
request_body=request_body,
)
validate_image_dimensions(
request_headers=request_headers,
request_body=request_body,
)
validate_image_file_size(
request_headers=request_headers,
request_body=request_body,
)
validate_max_num_results(
request_headers=request_headers,
request_body=request_body,
)
validate_include_target_data(
form = parse_multipart(
request_headers=request_headers,
request_body=request_body,
)
validate_extra_fields(form=form)
validate_image_field_given(form=form)
validate_image_is_image(form=form)
validate_image_format(form=form)
validate_image_dimensions(form=form)
validate_image_file_size(form=form)
validate_max_num_results(form=form)
validate_include_target_data(form=form)
return ValidatedQuery(database=database, form=form)
7 changes: 5 additions & 2 deletions src/mock_vws/_query_validators/auth_validators.py
Original file line number Diff line number Diff line change
Expand Up @@ -115,7 +115,7 @@ def validate_authorization(
request_body: bytes,
request_method: str,
databases: Iterable[CloudDatabase],
) -> None:
) -> CloudDatabase:
"""Validate the authorization header given to the query endpoint.

Args:
Expand All @@ -125,12 +125,15 @@ def validate_authorization(
request_method: The HTTP method of the request.
databases: All Vuforia databases.

Returns:
The database which the request's client keys belong to.

Raises:
AuthenticationFailureError: The "Authorization" header is not as
expected.
"""
try:
get_database_matching_client_keys(
return get_database_matching_client_keys(
request_headers=request_headers,
request_body=request_body,
request_method=request_method,
Expand Down
20 changes: 4 additions & 16 deletions src/mock_vws/_query_validators/fields_validators.py
Original file line number Diff line number Diff line change
@@ -1,38 +1,26 @@
"""Validators for the fields given."""

import logging
from collections.abc import Mapping

from beartype import beartype

from mock_vws._query_validators.exceptions import UnknownParametersError
from mock_vws._query_validators.multipart import parse_multipart
from mock_vws._query_validators.multipart import MultipartForm

_LOGGER = logging.getLogger(name=__name__)


@beartype
def validate_extra_fields(
*,
request_headers: Mapping[str, str],
request_body: bytes,
) -> None:
def validate_extra_fields(*, form: MultipartForm) -> None:
"""Validate that the no unknown fields are given.

Args:
request_headers: The headers sent with the request.
request_body: The body of the request.
form: The parsed body of the request.

Raises:
UnknownParametersError: Extra fields are given.
NoContentDispositionError: A part of the body has no
``Content-Disposition`` header.
"""
fields, files = parse_multipart(
request_headers=request_headers,
request_body=request_body,
)
parsed_keys = fields.keys() | files.keys()
parsed_keys = form.fields.keys() | form.files.keys()
known_parameters = {"image", "max_num_results", "include_target_data"}

if not parsed_keys - known_parameters:
Expand Down
Loading
Loading