190 lines
7.4 KiB
Python
190 lines
7.4 KiB
Python
class OverloadCandidate:
|
|
function: Function
|
|
mapped: list[MappedArgument]
|
|
|
|
class CallDispatcher(Generic[E]):
|
|
def get_result(
|
|
self,
|
|
location: Location,
|
|
callee: Type,
|
|
positional: list[TypedExpr[E]],
|
|
keywords: dict[str, TypedExpr[E]],
|
|
report_errors: bool = True,
|
|
) -> CallResult:
|
|
match callee:
|
|
case Function() as function:
|
|
valid: bool
|
|
mapped: list[MappedArgument[E]]
|
|
valid, mapped = self.map_call_arguments(
|
|
function, location, positional, keywords
|
|
)
|
|
valid = valid and self._are_arguments_valid(mapped, report_errors)
|
|
if not valid:
|
|
return CallResult(error=CallError.INVALID_ARGS)
|
|
return CallResult(result=function.returns)
|
|
case UnknownType():
|
|
return CallResult(result=UnknownType())
|
|
case DerivedType(type=base):
|
|
return self.get_result(
|
|
location, base, positional, keywords, report_errors
|
|
)
|
|
case AppliedType(body=body):
|
|
return self.get_result(
|
|
location, body, positional, keywords, report_errors
|
|
)
|
|
case OverloadedFunction(overloads=overloads):
|
|
res = self._match_overload(
|
|
overloads, location, positional, keywords, report_errors
|
|
)
|
|
if res[0] is None:
|
|
return CallResult(
|
|
error=CallError.NO_MATCHING_OVERLOAD,
|
|
message=res[1],
|
|
)
|
|
return CallResult(result=res[0].returns)
|
|
case GenericType():
|
|
unifier: Unifier = Unifier(self.types)
|
|
pos: list[Type] = [a[1] for a in positional]
|
|
kw: dict[str, Type] = {k: v[1] for k, v in keywords.items()}
|
|
unified: Optional[Type] = unifier.unify_call(callee, pos, kw)
|
|
if unified is None:
|
|
pos_str: str = ", ".join(str(t) for t in pos)
|
|
kw_str: str = ", ".join(f"{k}: {v}" for k, v in kw.items())
|
|
message: str = (
|
|
f"Could not unify {callee}={callee.body} with pos=[{pos_str}] and kw={{{kw_str}}}"
|
|
)
|
|
if report_errors:
|
|
self.reporter.error(location, message)
|
|
return CallResult(
|
|
error=CallError.IMPOSSIBLE_UNIFICATION,
|
|
message=message,
|
|
)
|
|
return self.get_result(
|
|
location,
|
|
unified,
|
|
positional,
|
|
keywords,
|
|
report_errors,
|
|
)
|
|
case _:
|
|
message: str = f"{callee} ({callee.__class__.__name__}) is not callable"
|
|
if report_errors:
|
|
self.reporter.error(location, message)
|
|
return CallResult(
|
|
error=CallError.NOT_CALLABLE,
|
|
message=message,
|
|
)
|
|
|
|
def _unwrap_function(
|
|
self,
|
|
callee: Type,
|
|
positional: list[TypedExpr[E]],
|
|
keywords: dict[str, TypedExpr[E]],
|
|
) -> Union[tuple[Function, None], tuple[None, CallError]]:
|
|
match callee:
|
|
case Function():
|
|
return callee, None
|
|
case DerivedType(type=base):
|
|
return self._unwrap_function(base, positional, keywords)
|
|
case AppliedType(body=body):
|
|
return self._unwrap_function(body, positional, keywords)
|
|
case GenericType():
|
|
unifier: Unifier = Unifier(self.types)
|
|
unified: Optional[Type] = unifier.unify_call(
|
|
callee,
|
|
[a[1] for a in positional],
|
|
{k: v[1] for k, v in keywords.items()},
|
|
)
|
|
if unified is None:
|
|
return None, CallError.IMPOSSIBLE_UNIFICATION
|
|
return self._unwrap_function(unified, positional, keywords)
|
|
case _:
|
|
return None, CallError.NOT_CALLABLE
|
|
|
|
def _match_overload(
|
|
self,
|
|
overloads: list[Type],
|
|
location: Location,
|
|
positional: list[TypedExpr[E]],
|
|
keywords: dict[str, TypedExpr[E]],
|
|
report_errors: bool = True,
|
|
) -> Union[tuple[Function, None], tuple[None, str]]:
|
|
candidates: list[OverloadCandidate] = []
|
|
errors: list[CallError] = []
|
|
for overload in overloads:
|
|
function, unwrap_error = self._unwrap_function(
|
|
overload, positional, keywords
|
|
)
|
|
if function is None:
|
|
errors.append(unwrap_error)
|
|
continue
|
|
|
|
valid, mapped = self.map_call_arguments(
|
|
function=function,
|
|
location=location,
|
|
positional=positional,
|
|
keywords=keywords,
|
|
report_errors=False,
|
|
)
|
|
if valid and self._are_arguments_valid(mapped, report_errors=False):
|
|
candidates.append(
|
|
OverloadCandidate(
|
|
function=function,
|
|
mapped=mapped,
|
|
)
|
|
)
|
|
|
|
pos_types: str = ", ".join(str(type) for _, type in positional)
|
|
kw_types: str = ", ".join(
|
|
f"{name}: {type}" for name, (_, type) in keywords.items()
|
|
)
|
|
for_args: str = f"for arguments pos=[{pos_types}] and kw={{{kw_types}}}"
|
|
|
|
n_candidates: int = len(candidates)
|
|
if n_candidates == 1:
|
|
return candidates[0].function, None
|
|
if n_candidates == 0:
|
|
overloads_str: str = ", ".join(map(str, overloads))
|
|
errors_str: str = ", ".join(errors)
|
|
message: str = (
|
|
f"No matching overload in [{overloads_str}] {for_args} (errors: {errors_str})"
|
|
)
|
|
if report_errors:
|
|
self.reporter.error(location, message)
|
|
return None, message
|
|
|
|
for i1, c1 in enumerate(candidates):
|
|
mapped1: list[MappedArgument[E]] = c1.mapped
|
|
best_match: bool = True
|
|
for i2, c2 in enumerate(candidates):
|
|
if i1 == i2:
|
|
continue
|
|
mapped2: list[MappedArgument[E]] = c2.mapped
|
|
if not self._are_mapped_subtypes(mapped1, mapped2):
|
|
best_match = False
|
|
break
|
|
self.logger.debug(f"{c1.function} is a full overload of {c2.function}")
|
|
if best_match:
|
|
return c1.function, None
|
|
|
|
candidates_str: str = ", ".join(
|
|
str(candidate.function) for candidate in candidates
|
|
)
|
|
message: str = f"Multiple matching overloads {for_args}: {candidates_str}"
|
|
if report_errors:
|
|
self.reporter.error(location, message)
|
|
return None, message
|
|
|
|
def _are_mapped_subtypes(
|
|
self, mapped1: list[MappedArgument[E]], mapped2: list[MappedArgument[E]]
|
|
) -> bool:
|
|
by_expr: dict[E, Type] = {}
|
|
for arg in mapped1:
|
|
by_expr[arg.arg_expr] = arg.parameter.type
|
|
|
|
for arg in mapped2:
|
|
type2: Type = arg.parameter.type
|
|
type1: Type = by_expr[arg.arg_expr]
|
|
if not self.types.is_subtype(type1, type2):
|
|
return False
|
|
return True |