141 lines
5.1 KiB
Python
141 lines
5.1 KiB
Python
import pathlib
|
|
import tempfile
|
|
|
|
import grpc_tools
|
|
from google.protobuf import descriptor
|
|
from google.protobuf import descriptor_pb2
|
|
from google.protobuf import descriptor_pool
|
|
from grpc_tools import protoc
|
|
|
|
|
|
SCRIPT_DIR = pathlib.Path(__file__).parent.resolve()
|
|
PROTO_DIR = SCRIPT_DIR.parent / 'solution' / 'proto'
|
|
PROTO_FILE = PROTO_DIR / 'messenger.proto'
|
|
WELL_KNOWN_PROTO_DIR = pathlib.Path(grpc_tools.__file__).parent / '_proto'
|
|
|
|
|
|
def compile_descriptor_set(output_path):
|
|
result = protoc.main([
|
|
'grpc_tools.protoc',
|
|
f'-I{PROTO_DIR}',
|
|
f'-I{WELL_KNOWN_PROTO_DIR}',
|
|
f'--descriptor_set_out={output_path}',
|
|
'--include_imports',
|
|
str(PROTO_FILE),
|
|
])
|
|
assert result == 0, 'messenger.proto must compile successfully'
|
|
|
|
descriptor_set = descriptor_pb2.FileDescriptorSet()
|
|
descriptor_set.ParseFromString(output_path.read_bytes())
|
|
return descriptor_set
|
|
|
|
|
|
def build_descriptor_pool(descriptor_set):
|
|
pool = descriptor_pool.DescriptorPool()
|
|
remaining = list(descriptor_set.file)
|
|
while remaining:
|
|
deferred = []
|
|
for file_descriptor in remaining:
|
|
try:
|
|
pool.Add(file_descriptor)
|
|
except TypeError:
|
|
deferred.append(file_descriptor)
|
|
assert len(deferred) < len(remaining), 'messenger.proto imports could not be resolved'
|
|
remaining = deferred
|
|
return pool
|
|
|
|
|
|
def require_singular_field(message_type, field_name, field_type, message_type_name=None):
|
|
assert field_name in message_type.fields_by_name, \
|
|
f'{message_type.full_name} must contain field {field_name}'
|
|
field = message_type.fields_by_name[field_name]
|
|
assert not field.is_repeated, f'{field.full_name} must be a singular field'
|
|
assert field.type == field_type, f'{field.full_name} has an invalid type'
|
|
if message_type_name is not None:
|
|
assert field.message_type is not None
|
|
assert field.message_type.full_name == message_type_name, \
|
|
f'{field.full_name} has an invalid message type'
|
|
|
|
|
|
def require_fields_can_coexist(message_type, field_names):
|
|
fields_by_oneof = {}
|
|
for field_name in field_names:
|
|
field = message_type.fields_by_name[field_name]
|
|
if field.containing_oneof is None:
|
|
continue
|
|
previous_field = fields_by_oneof.setdefault(field.containing_oneof.full_name, field_name)
|
|
assert previous_field == field_name, \
|
|
f'{message_type.full_name} fields must allow simultaneous values'
|
|
|
|
|
|
def test_proto_contract():
|
|
with tempfile.TemporaryDirectory() as temporary_directory:
|
|
descriptor_path = pathlib.Path(temporary_directory) / 'messenger.pb'
|
|
descriptor_set = compile_descriptor_set(descriptor_path)
|
|
|
|
submitted_file = next(
|
|
(file_descriptor for file_descriptor in descriptor_set.file
|
|
if pathlib.PurePosixPath(file_descriptor.name).name == PROTO_FILE.name),
|
|
None,
|
|
)
|
|
assert submitted_file is not None, 'messenger.proto descriptor is missing'
|
|
assert submitted_file.syntax == 'proto3', 'messenger.proto must use proto3 syntax'
|
|
assert submitted_file.package == 'mes_grpc', 'messenger.proto must use package mes_grpc'
|
|
|
|
pool = build_descriptor_pool(descriptor_set)
|
|
try:
|
|
messenger = pool.FindServiceByName('mes_grpc.MessengerServer')
|
|
except KeyError:
|
|
raise AssertionError('gRPC service must be named mes_grpc.MessengerServer') from None
|
|
|
|
assert 'SendMessage' in messenger.methods_by_name, \
|
|
'MessengerServer must contain method SendMessage'
|
|
send_message = messenger.methods_by_name['SendMessage']
|
|
assert not send_message.client_streaming and not send_message.server_streaming, \
|
|
'SendMessage must be unary'
|
|
require_singular_field(
|
|
send_message.input_type,
|
|
'author',
|
|
descriptor.FieldDescriptor.TYPE_STRING,
|
|
)
|
|
require_singular_field(
|
|
send_message.input_type,
|
|
'text',
|
|
descriptor.FieldDescriptor.TYPE_STRING,
|
|
)
|
|
require_fields_can_coexist(send_message.input_type, ('author', 'text'))
|
|
require_singular_field(
|
|
send_message.output_type,
|
|
'sendTime',
|
|
descriptor.FieldDescriptor.TYPE_MESSAGE,
|
|
'google.protobuf.Timestamp',
|
|
)
|
|
|
|
assert 'ReadMessages' in messenger.methods_by_name, \
|
|
'MessengerServer must contain method ReadMessages'
|
|
read_messages = messenger.methods_by_name['ReadMessages']
|
|
assert not read_messages.client_streaming and read_messages.server_streaming, \
|
|
'ReadMessages must be a unary request with a server stream response'
|
|
assert not read_messages.input_type.fields, \
|
|
'ReadMessages request must not contain fields'
|
|
require_singular_field(
|
|
read_messages.output_type,
|
|
'author',
|
|
descriptor.FieldDescriptor.TYPE_STRING,
|
|
)
|
|
require_singular_field(
|
|
read_messages.output_type,
|
|
'text',
|
|
descriptor.FieldDescriptor.TYPE_STRING,
|
|
)
|
|
require_singular_field(
|
|
read_messages.output_type,
|
|
'sendTime',
|
|
descriptor.FieldDescriptor.TYPE_MESSAGE,
|
|
'google.protobuf.Timestamp',
|
|
)
|
|
require_fields_can_coexist(
|
|
read_messages.output_type,
|
|
('author', 'text', 'sendTime'),
|
|
)
|