Files
newsolstice/src/typechecker/typechecker.hpp
2026-08-03 08:15:48 +10:00

94 lines
2.6 KiB
C++

#pragma once
#include "type.hpp"
#include "../parser/parser.hpp"
#include <string>
#include <unordered_map>
namespace Solstice {
enum class Narrowing {
None,
Left,
Right,
Both
};
class TypeChecker {
std::unordered_map<std::string, PossibleType> variables;
std::unordered_map<std::string, Type> types;
std::unordered_map<std::string, Function> functions;
std::unordered_map<TypePair, Type> addOverloads;
std::unordered_map<TypePair, Type> subtractOverloads;
std::unordered_map<TypePair, Type> multiplyOverloads;
std::unordered_map<TypePair, Type> divideOverloads;
std::unordered_map<TypePair, Type> equalOverloads;
std::unordered_map<TypePair, Type> notEqualOverloads;
std::unordered_map<TypePair, Type> greaterThanOverloads;
std::unordered_map<TypePair, Type> lesserThanOverloads;
Node* input;
bool inFunction = false;
void initOverloads();
void checkLiteralNodeType(Node& node);
void checkIdentifierNodeType(Node& node);
void checkTupleNodeType(Node& node);
void checkBindNodeType(Node& node);
void checkFunctionBindNodeType(Node& node);
void checkSetNodeType(Node& node);
void checkCodeBlockType(Node& node);
Narrowing doesNodeChildrenNeedNarrowing(Node& node);
void narrowBinaryNode(Node& node, std::unordered_map<TypePair, Type>& overloads);
void checkAddType(Node& node);
void checkSubtractType(Node& node);
void checkMultiplyType(Node& node);
void checkDivideType(Node& node);
void checkEqualType(Node& node);
void checkNotEqualType(Node& node);
void checkGreaterThanType(Node& node);
void checkLesserThanType(Node& node);
void checkNodeType(Node& node);
public:
TypeChecker() = delete;
TypeChecker(Node& node) : input(&node) {
initOverloads();
}
void setNode(Node& node) {
input = &node;
}
// Only for use as a public function.
void setVariable(const std::string& name, const Type& type) {
variables[name] = {{type}};
}
void setVariableUnknown(const std::string& name) {
variables[name].unknownType = true;
}
void checkTypes();
const Function* getFunction(const std::string& name) const {
auto it = functions.find(name);
if (it == functions.end()) {
return nullptr;
}
return &it->second;
}
};
}