| ... | ... | @@ -10,6 +10,12 @@ |
| 10 | 10 | #include "error.hpp" |
| 11 | 11 | #include "zig_llvm.hpp" |
| 12 | 12 | |
| 13 | struct BlockContext { |
| 14 | AstNode *node; |
| 15 | BlockContext *root; |
| 16 | BlockContext *parent; |
| 17 | }; |
| 18 | |
| 13 | 19 | static void add_node_error(CodeGen *g, AstNode *node, Buf *msg) { |
| 14 | 20 | g->errors.add_one(); |
| 15 | 21 | ErrorMsg *last_msg = &g->errors.last(); |
| ... | ... | @@ -229,6 +235,155 @@ static void preview_function_declarations(CodeGen *g, AstNode *node) { |
| 229 | 235 | } |
| 230 | 236 | } |
| 231 | 237 | |
| 238 | static TypeTableEntry * get_return_type(BlockContext *context) { |
| 239 | AstNode *fn_def_node = context->root->node; |
| 240 | assert(fn_def_node->type == NodeTypeFnDef); |
| 241 | AstNode *fn_proto_node = fn_def_node->data.fn_def.fn_proto; |
| 242 | assert(fn_proto_node->type == NodeTypeFnProto); |
| 243 | AstNode *return_type_node = fn_proto_node->data.fn_proto.return_type; |
| 244 | assert(return_type_node->codegen_node); |
| 245 | return return_type_node->codegen_node->data.type_node.entry; |
| 246 | } |
| 247 | |
| 248 | static void check_type_compatibility(CodeGen *g, AstNode *node, TypeTableEntry *expected_type, TypeTableEntry *actual_type) { |
| 249 | if (expected_type == actual_type) |
| 250 | return; // good |
| 251 | if (expected_type == g->builtin_types.entry_invalid || actual_type == g->builtin_types.entry_invalid) |
| 252 | return; // already complained |
| 253 | if (actual_type == g->builtin_types.entry_unreachable) |
| 254 | return; // TODO: is this true? |
| 255 | |
| 256 | // TODO better error message |
| 257 | add_node_error(g, node, buf_sprintf("type mismatch.")); |
| 258 | } |
| 259 | |
| 260 | static TypeTableEntry * analyze_expression(CodeGen *g, BlockContext *context, TypeTableEntry *expected_type, AstNode *node) { |
| 261 | switch (node->type) { |
| 262 | case NodeTypeBlock: |
| 263 | { |
| 264 | // TODO: nested block scopes |
| 265 | TypeTableEntry *return_type = g->builtin_types.entry_void; |
| 266 | for (int i = 0; i < node->data.block.statements.length; i += 1) { |
| 267 | AstNode *child = node->data.block.statements.at(i); |
| 268 | if (return_type == g->builtin_types.entry_unreachable) { |
| 269 | add_node_error(g, child, |
| 270 | buf_sprintf("unreachable code")); |
| 271 | break; |
| 272 | } |
| 273 | return_type = analyze_expression(g, context, nullptr, child); |
| 274 | } |
| 275 | return return_type; |
| 276 | } |
| 277 | |
| 278 | case NodeTypeReturnExpr: |
| 279 | { |
| 280 | TypeTableEntry *expected_return_type = get_return_type(context); |
| 281 | TypeTableEntry *actual_return_type; |
| 282 | if (node->data.return_expr.expr) { |
| 283 | actual_return_type = analyze_expression(g, context, expected_return_type, node->data.return_expr.expr); |
| 284 | } else { |
| 285 | actual_return_type = g->builtin_types.entry_void; |
| 286 | } |
| 287 | |
| 288 | if (actual_return_type == g->builtin_types.entry_unreachable) { |
| 289 | // "return exit(0)" should just be "exit(0)". |
| 290 | add_node_error(g, node, buf_sprintf("returning is unreachable.")); |
| 291 | actual_return_type = g->builtin_types.entry_invalid; |
| 292 | } |
| 293 | |
| 294 | check_type_compatibility(g, node, expected_return_type, actual_return_type); |
| 295 | return g->builtin_types.entry_unreachable; |
| 296 | } |
| 297 | |
| 298 | case NodeTypeBinOpExpr: |
| 299 | { |
| 300 | // TODO: think about expected types |
| 301 | analyze_expression(g, context, expected_type, node->data.bin_op_expr.op1); |
| 302 | analyze_expression(g, context, expected_type, node->data.bin_op_expr.op2); |
| 303 | return expected_type; |
| 304 | } |
| 305 | |
| 306 | case NodeTypeFnCallExpr: |
| 307 | { |
| 308 | Buf *name = hack_get_fn_call_name(g, node->data.fn_call_expr.fn_ref_expr); |
| 309 | |
| 310 | auto entry = g->fn_table.maybe_get(name); |
| 311 | if (!entry) { |
| 312 | add_node_error(g, node, |
| 313 | buf_sprintf("undefined function: '%s'", buf_ptr(name))); |
| 314 | // still analyze the parameters, even though we don't know what to expect |
| 315 | for (int i = 0; i < node->data.fn_call_expr.params.length; i += 1) { |
| 316 | AstNode *child = node->data.fn_call_expr.params.at(i); |
| 317 | analyze_expression(g, context, nullptr, child); |
| 318 | } |
| 319 | |
| 320 | return g->builtin_types.entry_invalid; |
| 321 | } else { |
| 322 | FnTableEntry *fn_table_entry = entry->value; |
| 323 | assert(fn_table_entry->proto_node->type == NodeTypeFnProto); |
| 324 | AstNodeFnProto *fn_proto = &fn_table_entry->proto_node->data.fn_proto; |
| 325 | |
| 326 | // count parameters |
| 327 | int expected_param_count = fn_proto->params.length; |
| 328 | int actual_param_count = node->data.fn_call_expr.params.length; |
| 329 | if (expected_param_count != actual_param_count) { |
| 330 | add_node_error(g, node, |
| 331 | buf_sprintf("wrong number of arguments. Expected %d, got %d.", |
| 332 | expected_param_count, actual_param_count)); |
| 333 | } |
| 334 | |
| 335 | // analyze each parameter |
| 336 | for (int i = 0; i < node->data.fn_call_expr.params.length; i += 1) { |
| 337 | AstNode *child = node->data.fn_call_expr.params.at(i); |
| 338 | // determine the expected type for each parameter |
| 339 | TypeTableEntry *expected_param_type = nullptr; |
| 340 | if (i < fn_proto->params.length) { |
| 341 | AstNode *param_decl_node = fn_proto->params.at(i); |
| 342 | assert(param_decl_node->type == NodeTypeParamDecl); |
| 343 | AstNode *param_type_node = param_decl_node->data.param_decl.type; |
| 344 | if (param_type_node->codegen_node) |
| 345 | expected_param_type = param_type_node->codegen_node->data.type_node.entry; |
| 346 | } |
| 347 | analyze_expression(g, context, expected_param_type, child); |
| 348 | } |
| 349 | |
| 350 | TypeTableEntry *return_type = fn_proto->return_type->codegen_node->data.type_node.entry; |
| 351 | check_type_compatibility(g, node, expected_type, return_type); |
| 352 | return return_type; |
| 353 | } |
| 354 | } |
| 355 | |
| 356 | case NodeTypeNumberLiteral: |
| 357 | // TODO: generic literal int type |
| 358 | return g->builtin_types.entry_i32; |
| 359 | |
| 360 | case NodeTypeStringLiteral: |
| 361 | zig_panic("TODO"); |
| 362 | |
| 363 | case NodeTypeUnreachable: |
| 364 | return g->builtin_types.entry_unreachable; |
| 365 | |
| 366 | case NodeTypeSymbol: |
| 367 | // look up symbol in symbol table |
| 368 | zig_panic("TODO"); |
| 369 | |
| 370 | case NodeTypeCastExpr: |
| 371 | case NodeTypePrefixOpExpr: |
| 372 | zig_panic("TODO"); |
| 373 | case NodeTypeDirective: |
| 374 | case NodeTypeFnDecl: |
| 375 | case NodeTypeFnProto: |
| 376 | case NodeTypeParamDecl: |
| 377 | case NodeTypeType: |
| 378 | case NodeTypeRoot: |
| 379 | case NodeTypeRootExportDecl: |
| 380 | case NodeTypeExternBlock: |
| 381 | case NodeTypeFnDef: |
| 382 | zig_unreachable(); |
| 383 | } |
| 384 | zig_unreachable(); |
| 385 | } |
| 386 | |
| 232 | 387 | static void check_fn_def_control_flow(CodeGen *g, AstNode *node) { |
| 233 | 388 | // Follow the execution flow and make sure the code returns appropriately. |
| 234 | 389 | // * A `return` statement in an unreachable type function should be an error. |
| ... | ... | @@ -282,74 +437,6 @@ static void check_fn_def_control_flow(CodeGen *g, AstNode *node) { |
| 282 | 437 | } |
| 283 | 438 | } |
| 284 | 439 | |
| 285 | | static void analyze_expression(CodeGen *g, AstNode *node) { |
| 286 | | switch (node->type) { |
| 287 | | case NodeTypeBlock: |
| 288 | | for (int i = 0; i < node->data.block.statements.length; i += 1) { |
| 289 | | AstNode *child = node->data.block.statements.at(i); |
| 290 | | analyze_expression(g, child); |
| 291 | | } |
| 292 | | break; |
| 293 | | case NodeTypeReturnExpr: |
| 294 | | if (node->data.return_expr.expr) { |
| 295 | | analyze_expression(g, node->data.return_expr.expr); |
| 296 | | } |
| 297 | | break; |
| 298 | | case NodeTypeBinOpExpr: |
| 299 | | analyze_expression(g, node->data.bin_op_expr.op1); |
| 300 | | analyze_expression(g, node->data.bin_op_expr.op2); |
| 301 | | break; |
| 302 | | case NodeTypeFnCallExpr: |
| 303 | | { |
| 304 | | Buf *name = hack_get_fn_call_name(g, node->data.fn_call_expr.fn_ref_expr); |
| 305 | | |
| 306 | | auto entry = g->fn_table.maybe_get(name); |
| 307 | | if (!entry) { |
| 308 | | add_node_error(g, node, |
| 309 | | buf_sprintf("undefined function: '%s'", buf_ptr(name))); |
| 310 | | } else { |
| 311 | | FnTableEntry *fn_table_entry = entry->value; |
| 312 | | assert(fn_table_entry->proto_node->type == NodeTypeFnProto); |
| 313 | | int expected_param_count = fn_table_entry->proto_node->data.fn_proto.params.length; |
| 314 | | int actual_param_count = node->data.fn_call_expr.params.length; |
| 315 | | if (expected_param_count != actual_param_count) { |
| 316 | | add_node_error(g, node, |
| 317 | | buf_sprintf("wrong number of arguments. Expected %d, got %d.", |
| 318 | | expected_param_count, actual_param_count)); |
| 319 | | } |
| 320 | | } |
| 321 | | |
| 322 | | for (int i = 0; i < node->data.fn_call_expr.params.length; i += 1) { |
| 323 | | AstNode *child = node->data.fn_call_expr.params.at(i); |
| 324 | | analyze_expression(g, child); |
| 325 | | } |
| 326 | | break; |
| 327 | | } |
| 328 | | case NodeTypeCastExpr: |
| 329 | | zig_panic("TODO"); |
| 330 | | break; |
| 331 | | case NodeTypePrefixOpExpr: |
| 332 | | zig_panic("TODO"); |
| 333 | | break; |
| 334 | | case NodeTypeNumberLiteral: |
| 335 | | case NodeTypeStringLiteral: |
| 336 | | case NodeTypeUnreachable: |
| 337 | | case NodeTypeSymbol: |
| 338 | | // nothing to do |
| 339 | | break; |
| 340 | | case NodeTypeDirective: |
| 341 | | case NodeTypeFnDecl: |
| 342 | | case NodeTypeFnProto: |
| 343 | | case NodeTypeParamDecl: |
| 344 | | case NodeTypeType: |
| 345 | | case NodeTypeRoot: |
| 346 | | case NodeTypeRootExportDecl: |
| 347 | | case NodeTypeExternBlock: |
| 348 | | case NodeTypeFnDef: |
| 349 | | zig_unreachable(); |
| 350 | | } |
| 351 | | } |
| 352 | | |
| 353 | 440 | static void analyze_top_level_declaration(CodeGen *g, AstNode *node) { |
| 354 | 441 | switch (node->type) { |
| 355 | 442 | case NodeTypeFnDef: |
| ... | ... | @@ -371,7 +458,13 @@ static void analyze_top_level_declaration(CodeGen *g, AstNode *node) { |
| 371 | 458 | } |
| 372 | 459 | |
| 373 | 460 | check_fn_def_control_flow(g, node); |
| 374 | | analyze_expression(g, node->data.fn_def.body); |
| 461 | |
| 462 | BlockContext context; |
| 463 | context.node = node; |
| 464 | context.root = &context; |
| 465 | context.parent = nullptr; |
| 466 | TypeTableEntry *expected_type = fn_proto->return_type->codegen_node->data.type_node.entry; |
| 467 | analyze_expression(g, &context, expected_type, node->data.fn_def.body); |
| 375 | 468 | } |
| 376 | 469 | break; |
| 377 | 470 | |
| ... | ... | @@ -424,6 +517,12 @@ static void analyze_root(CodeGen *g, AstNode *node) { |
| 424 | 517 | } |
| 425 | 518 | |
| 426 | 519 | static void define_primitive_types(CodeGen *g) { |
| 520 | { |
| 521 | // if this type is anywhere in the AST, we should never hit codegen. |
| 522 | TypeTableEntry *entry = allocate<TypeTableEntry>(1); |
| 523 | buf_init_from_str(&entry->name, "(invalid)"); |
| 524 | g->builtin_types.entry_invalid = entry; |
| 525 | } |
| 427 | 526 | { |
| 428 | 527 | TypeTableEntry *entry = allocate<TypeTableEntry>(1); |
| 429 | 528 | entry->type_ref = LLVMInt8Type(); |
| ... | ... | @@ -450,15 +549,12 @@ static void define_primitive_types(CodeGen *g) { |
| 450 | 549 | LLVMZigEncoding_DW_ATE_unsigned()); |
| 451 | 550 | g->type_table.put(&entry->name, entry); |
| 452 | 551 | g->builtin_types.entry_void = entry; |
| 453 | | |
| 454 | | // invalid types are void |
| 455 | | g->builtin_types.entry_invalid = entry; |
| 456 | 552 | } |
| 457 | 553 | { |
| 458 | 554 | TypeTableEntry *entry = allocate<TypeTableEntry>(1); |
| 459 | 555 | entry->type_ref = LLVMVoidType(); |
| 460 | 556 | buf_init_from_str(&entry->name, "unreachable"); |
| 461 | | entry->di_type = g->builtin_types.entry_invalid->di_type; |
| 557 | entry->di_type = g->builtin_types.entry_void->di_type; |
| 462 | 558 | g->type_table.put(&entry->name, entry); |
| 463 | 559 | g->builtin_types.entry_unreachable = entry; |
| 464 | 560 | } |