rlsl 0.1.1 → 1.0.1

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (91) hide show
  1. checksums.yaml +4 -4
  2. data/.rubocop.yml +8 -0
  3. data/CHANGELOG.md +13 -2
  4. data/README.md +101 -26
  5. data/Rakefile +7 -0
  6. data/lib/rlsl/base_translator/call_parser.rb +79 -0
  7. data/lib/rlsl/base_translator/code_rewriter.rb +107 -0
  8. data/lib/rlsl/base_translator/code_scanner.rb +140 -0
  9. data/lib/rlsl/base_translator.rb +182 -64
  10. data/lib/rlsl/code_generator/math_prelude.rb +172 -0
  11. data/lib/rlsl/code_generator/ruby_wrapper_generator.rb +137 -0
  12. data/lib/rlsl/code_generator/shader_function_generator.rb +19 -0
  13. data/lib/rlsl/code_generator/template_context.rb +44 -0
  14. data/lib/rlsl/code_generator/uniform_struct_generator.rb +29 -0
  15. data/lib/rlsl/code_generator.rb +30 -202
  16. data/lib/rlsl/compiled_shader.rb +5 -13
  17. data/lib/rlsl/errors.rb +29 -0
  18. data/lib/rlsl/function_context.rb +31 -14
  19. data/lib/rlsl/glsl/translator.rb +35 -41
  20. data/lib/rlsl/msl/shader.rb +34 -46
  21. data/lib/rlsl/msl/translator.rb +38 -34
  22. data/lib/rlsl/msl/uniform_buffer_packer.rb +72 -0
  23. data/lib/rlsl/prism/ast_visitor/control_flow_visiting.rb +137 -0
  24. data/lib/rlsl/prism/ast_visitor/definition_visiting.rb +92 -0
  25. data/lib/rlsl/prism/ast_visitor/expression_visiting.rb +172 -0
  26. data/lib/rlsl/prism/ast_visitor/scope_context.rb +48 -0
  27. data/lib/rlsl/prism/ast_visitor/visitor_registry.rb +21 -0
  28. data/lib/rlsl/prism/ast_visitor.rb +100 -286
  29. data/lib/rlsl/prism/builtins/function_registry.rb +131 -0
  30. data/lib/rlsl/prism/builtins/operator_rules.rb +123 -0
  31. data/lib/rlsl/prism/builtins/swizzle_rules.rb +47 -0
  32. data/lib/rlsl/prism/builtins.rb +39 -148
  33. data/lib/rlsl/prism/compilation_unit.rb +7 -0
  34. data/lib/rlsl/prism/emitters/base_emitter/control_flow_emission.rb +126 -0
  35. data/lib/rlsl/prism/emitters/base_emitter/definition_emission.rb +134 -0
  36. data/lib/rlsl/prism/emitters/base_emitter/expression_emission.rb +108 -0
  37. data/lib/rlsl/prism/emitters/base_emitter/statement_emission.rb +95 -0
  38. data/lib/rlsl/prism/emitters/base_emitter.rb +120 -414
  39. data/lib/rlsl/prism/emitters/c_emitter.rb +165 -112
  40. data/lib/rlsl/prism/emitters/glsl_emitter.rb +63 -50
  41. data/lib/rlsl/prism/emitters/msl_emitter.rb +67 -52
  42. data/lib/rlsl/prism/emitters/target_emitter.rb +124 -0
  43. data/lib/rlsl/prism/emitters/target_profile.rb +34 -0
  44. data/lib/rlsl/prism/emitters/wgsl_emitter.rb +217 -58
  45. data/lib/rlsl/prism/errors.rb +9 -0
  46. data/lib/rlsl/prism/ir/control_flow.rb +85 -0
  47. data/lib/rlsl/prism/ir/definitions.rb +82 -0
  48. data/lib/rlsl/prism/ir/expressions.rb +201 -0
  49. data/lib/rlsl/prism/ir/node.rb +21 -0
  50. data/lib/rlsl/prism/ir/nodes.rb +4 -371
  51. data/lib/rlsl/prism/ir/traversal.rb +66 -0
  52. data/lib/rlsl/prism/mutation_analyzer.rb +30 -0
  53. data/lib/rlsl/prism/node_traversal.rb +41 -0
  54. data/lib/rlsl/prism/parameter_list.rb +46 -0
  55. data/lib/rlsl/prism/return_flow_validator.rb +73 -0
  56. data/lib/rlsl/prism/source_extractor/block_locator.rb +50 -0
  57. data/lib/rlsl/prism/source_extractor.rb +19 -137
  58. data/lib/rlsl/prism/source_unit/parser.rb +88 -0
  59. data/lib/rlsl/prism/source_unit.rb +42 -0
  60. data/lib/rlsl/prism/target_capability_validator.rb +110 -0
  61. data/lib/rlsl/prism/transpiler.rb +99 -59
  62. data/lib/rlsl/prism/type_inference/call_type_resolver.rb +59 -0
  63. data/lib/rlsl/prism/type_inference/call_validator.rb +75 -0
  64. data/lib/rlsl/prism/type_inference/collection_type_resolver.rb +106 -0
  65. data/lib/rlsl/prism/type_inference/control_flow_inferer.rb +147 -0
  66. data/lib/rlsl/prism/type_inference/definition_inferer.rb +56 -0
  67. data/lib/rlsl/prism/type_inference/expression_inferer.rb +96 -0
  68. data/lib/rlsl/prism/type_inference/field_type_resolver.rb +38 -0
  69. data/lib/rlsl/prism/type_inference/inferer_registry.rb +38 -0
  70. data/lib/rlsl/prism/type_inference/scope_stack.rb +47 -0
  71. data/lib/rlsl/prism/type_inference/type_environment.rb +112 -0
  72. data/lib/rlsl/prism/type_inference/type_shapes.rb +33 -0
  73. data/lib/rlsl/prism/type_inference.rb +120 -249
  74. data/lib/rlsl/runtime_shader.rb +47 -0
  75. data/lib/rlsl/shader_builder/build_service.rb +104 -0
  76. data/lib/rlsl/shader_builder/native_extension_compiler.rb +97 -0
  77. data/lib/rlsl/shader_builder/shader_definition.rb +68 -0
  78. data/lib/rlsl/shader_builder/source_resolver.rb +109 -0
  79. data/lib/rlsl/shader_builder.rb +60 -111
  80. data/lib/rlsl/shader_name.rb +18 -0
  81. data/lib/rlsl/types/catalog.rb +47 -0
  82. data/lib/rlsl/types/target_resolver.rb +15 -0
  83. data/lib/rlsl/types/type_spec.rb +167 -0
  84. data/lib/rlsl/types/value_normalizer.rb +86 -0
  85. data/lib/rlsl/types.rb +9 -31
  86. data/lib/rlsl/uniform_context.rb +22 -11
  87. data/lib/rlsl/version.rb +1 -1
  88. data/lib/rlsl/wgsl/translator.rb +46 -39
  89. data/lib/rlsl/wgsl/uniform_layout.rb +25 -0
  90. data/lib/rlsl.rb +38 -15
  91. metadata +76 -11
@@ -3,11 +3,17 @@
3
3
  require "prism"
4
4
 
5
5
  require_relative "ir/nodes"
6
+ require_relative "compilation_unit"
7
+ require_relative "source_unit"
6
8
  require_relative "source_extractor"
7
9
  require_relative "builtins"
8
10
  require_relative "ast_visitor"
9
11
  require_relative "type_inference"
12
+ require_relative "target_capability_validator"
13
+ require_relative "mutation_analyzer"
14
+ require_relative "return_flow_validator"
10
15
  require_relative "emitters/base_emitter"
16
+ require_relative "emitters/target_emitter"
11
17
  require_relative "emitters/c_emitter"
12
18
  require_relative "emitters/msl_emitter"
13
19
  require_relative "emitters/wgsl_emitter"
@@ -23,70 +29,115 @@ module RLSL
23
29
  glsl: Emitters::GLSLEmitter
24
30
  }.freeze
25
31
 
26
- attr_reader :ir, :uniforms, :custom_functions
32
+ attr_reader :uniforms, :custom_functions
27
33
 
28
- def initialize(uniforms = {}, custom_functions = {})
34
+ def initialize(uniforms = {}, custom_functions = {}, globals: {})
29
35
  @uniforms = uniforms
30
36
  @custom_functions = custom_functions
37
+ @globals = globals
31
38
  @source_extractor = SourceExtractor.new
32
- @ir = nil
33
39
  end
34
40
 
35
- def parse_block(block)
36
- source = @source_extractor.extract(block)
37
- parse_source(source)
41
+ def compile_block(block)
42
+ compile_unit(@source_extractor.extract_unit(block))
38
43
  end
39
44
 
40
- def parse_source(source)
41
- params, body = extract_block_body(source)
42
-
43
- visitor = ASTVisitor.new(uniforms: @uniforms, params: params)
44
- @ir = visitor.parse(body)
45
-
46
- inference = TypeInference.new(@uniforms, @custom_functions)
47
- inference.register(:frag_coord, :vec2)
48
- inference.register(:resolution, :vec2)
49
- inference.infer(@ir)
50
-
51
- @ir
45
+ def compile_source(source)
46
+ compile_unit(source_unit(source))
52
47
  end
53
48
 
54
- def emit(target, needs_return: true)
55
- raise "No IR parsed yet. Call parse_block or parse_source first." unless @ir
49
+ def compile_helpers(block, function_signatures = {})
50
+ compile_unit(
51
+ @source_extractor.extract_unit(block).without_params,
52
+ function_signatures: function_signatures
53
+ )
54
+ end
56
55
 
57
- emitter_class = TARGETS[target.to_sym]
58
- raise "Unknown target: #{target}" unless emitter_class
56
+ def compile_helpers_source(source, function_signatures = {})
57
+ compile_unit(
58
+ source_unit(source).without_params,
59
+ function_signatures: function_signatures
60
+ )
61
+ end
59
62
 
60
- emitter = emitter_class.new
61
- emitter.emit(@ir, needs_return: needs_return)
63
+ def emit(target, compilation:, needs_return: true)
64
+ emitter = resolve_emitter(target)
65
+ validate_target_capabilities!(compilation.ir, target)
66
+ ReturnFlowValidator.new.validate!(compilation.ir, needs_return: needs_return)
67
+ emitter.emit(compilation.ir, needs_return: needs_return)
62
68
  end
63
69
 
64
70
  def transpile(block, target)
65
- parse_block(block)
66
- emit(target)
71
+ emit(target, compilation: compile_block(block))
67
72
  end
68
73
 
69
74
  def transpile_source(source, target)
70
- parse_source(source)
71
- emit(target)
75
+ emit(target, compilation: compile_source(source))
72
76
  end
73
77
 
74
78
  def transpile_helpers(block, target, function_signatures = {})
75
- source = @source_extractor.extract(block)
76
- _, body = extract_block_body(source)
79
+ emit(
80
+ target,
81
+ needs_return: false,
82
+ compilation: compile_helpers(block, function_signatures)
83
+ )
84
+ end
77
85
 
78
- visitor = ASTVisitor.new(uniforms: @uniforms)
79
- @ir = visitor.parse(body)
86
+ def transpile_helpers_source(source, target, function_signatures = {})
87
+ emit(
88
+ target,
89
+ needs_return: false,
90
+ compilation: compile_helpers_source(source, function_signatures)
91
+ )
92
+ end
80
93
 
81
- apply_function_signatures(@ir, function_signatures)
94
+ private
82
95
 
83
- inference = TypeInference.new(@uniforms, @custom_functions)
84
- inference.infer(@ir)
96
+ def source_unit(source)
97
+ return source if source.is_a?(SourceUnit)
85
98
 
86
- emit(target, needs_return: false)
99
+ SourceUnit.from_source(source)
87
100
  end
88
101
 
89
- private
102
+ def build_ir(unit)
103
+ visitor = ASTVisitor.new(
104
+ uniforms: @uniforms,
105
+ params: unit.params,
106
+ source_name: unit.source_name,
107
+ line_offset: unit.line_offset
108
+ )
109
+ visitor.parse(unit.body)
110
+ end
111
+
112
+ def compile_unit(unit, function_signatures: nil)
113
+ ir = build_ir(unit)
114
+ apply_function_signatures(ir, function_signatures || {})
115
+ MutationAnalyzer.new.analyze(ir)
116
+ infer_ir(ir)
117
+ CompilationUnit.new(source_unit: unit, ir: ir)
118
+ end
119
+
120
+ def infer_ir(ir)
121
+ inference = TypeInference.new(@uniforms, @custom_functions, globals: @globals)
122
+ register_pipeline_symbols(inference)
123
+ inference.infer(ir)
124
+ end
125
+
126
+ def register_pipeline_symbols(inference)
127
+ inference.register(:frag_coord, :vec2)
128
+ inference.register(:resolution, :vec2)
129
+ end
130
+
131
+ def resolve_emitter(target)
132
+ emitter_class = TARGETS[target.to_sym]
133
+ raise RLSL::Error, "Unknown target: #{target}" unless emitter_class
134
+
135
+ emitter_class.new
136
+ end
137
+
138
+ def validate_target_capabilities!(ir, target)
139
+ TargetCapabilityValidator.new.validate!(ir, target)
140
+ end
90
141
 
91
142
  def apply_function_signatures(ir, signatures)
92
143
  return unless ir.is_a?(IR::Block)
@@ -95,33 +146,22 @@ module RLSL
95
146
  next unless stmt.is_a?(IR::FunctionDefinition)
96
147
 
97
148
  sig = signatures[stmt.name]
98
- next unless sig
149
+ unless sig
150
+ raise SignatureError.new(
151
+ "Function #{stmt.name} requires an explicit signature in functions"
152
+ ).with_source_location(stmt.location)
153
+ end
99
154
 
100
155
  stmt.return_type = sig[:returns]
101
156
  stmt.param_types = sig[:params] || {}
102
- end
103
- end
104
-
105
- def extract_block_body(source)
106
- lines = source.strip.lines
107
- params = []
108
-
109
- first_line = lines.first&.strip || ""
110
-
111
- if first_line.start_with?("|")
112
- param_end = first_line.index("|", 1)
113
- if param_end
114
- param_str = first_line[1...param_end]
115
- params = param_str.split(",").map { |p| p.strip.to_sym }
116
- lines[0] = first_line[(param_end + 1)..]
157
+ unless stmt.params == stmt.param_types.keys
158
+ error = SignatureError.new(
159
+ "Function #{stmt.name} parameters #{stmt.params.inspect} " \
160
+ "do not match signature #{stmt.param_types.keys.inspect}"
161
+ )
162
+ raise error.with_source_location(stmt.location)
117
163
  end
118
164
  end
119
-
120
- lines.shift while lines.first&.strip&.empty?
121
- lines.pop while lines.last&.strip&.empty?
122
-
123
- body = lines.join.strip
124
- [params, body]
125
165
  end
126
166
  end
127
167
  end
@@ -0,0 +1,59 @@
1
+ # frozen_string_literal: true
2
+
3
+ module RLSL
4
+ module Prism
5
+ class CallTypeResolver
6
+ def initialize(custom_functions:, call_validator:)
7
+ @custom_functions = custom_functions
8
+ @call_validator = call_validator
9
+ end
10
+
11
+ def resolve(node)
12
+ sig = Builtins.function_signature(node.name)
13
+ return resolve_builtin(node, sig) if sig
14
+
15
+ custom_function = @custom_functions[node.name.to_sym]
16
+ return resolve_custom(node, custom_function) if custom_function
17
+
18
+ raise SignatureError, "Unknown shader function #{node.name}"
19
+ end
20
+
21
+ private
22
+
23
+ def resolve_builtin(node, signature)
24
+ arg_types = effective_arg_types(node)
25
+ @call_validator.validate_builtin!(node, arg_types, signature)
26
+ return_type = Builtins.resolve_return_type(signature[:returns], arg_types)
27
+ unless return_type
28
+ raise SignatureError, "Incompatible argument types for #{node.name}: #{arg_types.join(', ')}"
29
+ end
30
+
31
+ node.expected_arg_types = expected_builtin_types(signature, arg_types.length, return_type)
32
+ return_type
33
+ end
34
+
35
+ def resolve_custom(node, signature)
36
+ @call_validator.validate_custom!(node.name, effective_arg_types(node), signature)
37
+ node.expected_arg_types = signature[:params]&.values || []
38
+ signature[:returns]
39
+ end
40
+
41
+ def effective_arg_types(node)
42
+ types = node.args.map(&:type)
43
+ node.receiver ? [node.receiver.type, *types] : types
44
+ end
45
+
46
+ def expected_builtin_types(signature, argument_count, return_type)
47
+ expected = signature[:args].first(argument_count)
48
+ case signature[:returns]
49
+ when :common, :floating
50
+ expected.map { |type| type == :any ? return_type : type }
51
+ when :interpolated
52
+ expected.each_with_index.map { |type, index| index < 2 && type == :any ? return_type : type }
53
+ else
54
+ expected
55
+ end
56
+ end
57
+ end
58
+ end
59
+ end
@@ -0,0 +1,75 @@
1
+ # frozen_string_literal: true
2
+
3
+ module RLSL
4
+ module Prism
5
+ class CallValidator
6
+ def validate_builtin!(node, arg_types, signature)
7
+ validate_signature!(
8
+ node.name,
9
+ arg_types,
10
+ signature[:args],
11
+ variadic: signature[:variadic],
12
+ min_args: signature[:min_args]
13
+ ) unless skip_builtin_validation?(node, signature)
14
+ end
15
+
16
+ def validate_custom!(name, arg_types, signature)
17
+ params = signature[:params]
18
+ unless params
19
+ return if arg_types.empty?
20
+
21
+ raise SignatureError, "Function #{name} requires explicit parameter types"
22
+ end
23
+
24
+ validate_signature!(name, arg_types, params.values)
25
+ end
26
+
27
+ private
28
+
29
+ def validate_signature!(name, arg_types, expected_types, variadic: false, min_args: nil)
30
+ validate_argument_count!(name, arg_types.length, expected_types.length, variadic: variadic, min_args: min_args)
31
+
32
+ arg_types.each_with_index do |actual_type, index|
33
+ expected_type = expected_types[index]
34
+ next if compatible_argument_type?(expected_type, actual_type)
35
+
36
+ raise SignatureError,
37
+ "Invalid argument #{index + 1} for #{name}: expected #{expected_type}, got #{actual_type || :unknown}"
38
+ end
39
+ end
40
+
41
+ def validate_argument_count!(name, actual_count, expected_count, variadic:, min_args:)
42
+ return if valid_argument_count?(actual_count, expected_count, variadic: variadic, min_args: min_args)
43
+
44
+ raise SignatureError,
45
+ "Wrong number of arguments for #{name}: expected #{expected_count_description(expected_count, variadic, min_args)}, got #{actual_count}"
46
+ end
47
+
48
+ def valid_argument_count?(actual_count, expected_count, variadic:, min_args:)
49
+ return actual_count == expected_count unless variadic
50
+
51
+ minimum = min_args || expected_count
52
+ actual_count.between?(minimum, expected_count)
53
+ end
54
+
55
+ def expected_count_description(expected_count, variadic, min_args)
56
+ return expected_count.to_s unless variadic
57
+
58
+ minimum = min_args || expected_count
59
+ minimum == expected_count ? minimum.to_s : "#{minimum}..#{expected_count}"
60
+ end
61
+
62
+ def compatible_argument_type?(expected_type, actual_type)
63
+ return true if expected_type == :any
64
+ return true if expected_type == actual_type
65
+ return true if expected_type == :float && actual_type == :int
66
+
67
+ false
68
+ end
69
+
70
+ def skip_builtin_validation?(node, signature)
71
+ signature[:variadic] && node.args.empty? && !node.type.nil?
72
+ end
73
+ end
74
+ end
75
+ end
@@ -0,0 +1,106 @@
1
+ # frozen_string_literal: true
2
+
3
+ module RLSL
4
+ module Prism
5
+ class CollectionTypeResolver
6
+ include TypeShapes
7
+
8
+ def initialize(type_environment:, custom_functions:, register:)
9
+ @type_environment = type_environment
10
+ @custom_functions = custom_functions
11
+ @register = register
12
+ end
13
+
14
+ def resolve_array_literal(node)
15
+ element_types = node.elements.map(&:type)
16
+ element_type = element_types.empty? ? :float : Builtins.common_type(element_types)
17
+ unless element_type
18
+ raise SignatureError, "Array elements have incompatible types: #{element_types.uniq.join(', ')}"
19
+ end
20
+
21
+ TypeShapes.array(element_type, node.elements.length)
22
+ end
23
+
24
+ def resolve_array_index(node)
25
+ array_type = node.array.type
26
+ return TypeShapes.element_type(array_type) if TypeShapes.array?(array_type)
27
+ return @type_environment.array_element_type(node.array.name) || :float if node.array.is_a?(IR::VarRef)
28
+
29
+ :float
30
+ end
31
+
32
+ def resolve_global_decl(node)
33
+ if node.initializer.is_a?(IR::ArrayLiteral)
34
+ node.array_size ||= node.initializer.elements.length
35
+ node.element_type ||= TypeShapes.element_type(node.initializer.type)
36
+ return TypeShapes.array(node.element_type, node.array_size)
37
+ end
38
+
39
+ node.initializer&.type
40
+ end
41
+
42
+ def assign_multiple_targets(node)
43
+ value_type = node.value.type
44
+ return assign_tuple_targets(node, node.value.elements.map(&:type)) if node.value.is_a?(IR::ArrayLiteral)
45
+ return assign_tuple_targets(node, value_type.types) if value_type.is_a?(IR::TupleType)
46
+ return assign_array_targets(node, TypeShapes.element_type(value_type)) if TypeShapes.array?(value_type)
47
+ return assign_custom_targets(node, @custom_functions[node.value.name]) if custom_multi_return?(node.value)
48
+
49
+ nil
50
+ end
51
+
52
+ private
53
+
54
+ def assign_tuple_targets(node, types)
55
+ validate_target_count!(node.targets, types.length)
56
+ node.targets.each_with_index do |target, index|
57
+ assign_target(node, target, index, types[index])
58
+ end
59
+ end
60
+
61
+ def assign_array_targets(node, type)
62
+ count = node.value.type.element_count
63
+ validate_target_count!(node.targets, count) if count
64
+ node.targets.each_with_index do |target, index|
65
+ assign_target(node, target, index, type)
66
+ end
67
+ end
68
+
69
+ def assign_custom_targets(node, signature)
70
+ returns = signature[:returns]
71
+ return unless returns.is_a?(Array)
72
+
73
+ assign_tuple_targets(node, returns)
74
+ end
75
+
76
+ def assign_target(node, target, index, type)
77
+ unless node.declarations[index]
78
+ existing_type = @type_environment.lookup(target.name)
79
+ unless compatible_assignment?(existing_type, type)
80
+ raise SignatureError, "Cannot assign #{type} to #{target.name} (#{existing_type})"
81
+ end
82
+
83
+ target.type = existing_type
84
+ return
85
+ end
86
+
87
+ target.type = type
88
+ @register.call(target.name, type)
89
+ end
90
+
91
+ def validate_target_count!(targets, value_count)
92
+ return if targets.length == value_count
93
+
94
+ raise SignatureError, "Multiple assignment has #{targets.length} targets for #{value_count} values"
95
+ end
96
+
97
+ def compatible_assignment?(target_type, value_type)
98
+ target_type == value_type || (target_type == :float && value_type == :int)
99
+ end
100
+
101
+ def custom_multi_return?(value)
102
+ value.is_a?(IR::FuncCall) && @custom_functions.key?(value.name)
103
+ end
104
+ end
105
+ end
106
+ end
@@ -0,0 +1,147 @@
1
+ # frozen_string_literal: true
2
+
3
+ module RLSL
4
+ module Prism
5
+ class ControlFlowInferer
6
+ def initialize(infer:, infer_in_scope:, lookup:)
7
+ @infer = infer
8
+ @infer_in_scope = infer_in_scope
9
+ @lookup = lookup
10
+ end
11
+
12
+ def infer_if_statement(node)
13
+ @infer.call(node.condition)
14
+ if node.hoisted_variables.empty?
15
+ @infer.call(node.then_branch, scoped: true)
16
+ @infer.call(node.else_branch, scoped: true) if node.else_branch
17
+ else
18
+ @infer.call(node.then_branch)
19
+ @infer.call(node.else_branch) if node.else_branch
20
+ end
21
+
22
+ node.hoisted_variables.each_key do |name|
23
+ node.hoisted_variables[name] = @lookup.call(name)
24
+ end
25
+
26
+ node.type = node.then_branch&.type
27
+ node
28
+ end
29
+
30
+ def infer_ternary(node)
31
+ @infer.call(node.condition)
32
+ @infer.call(node.then_expr)
33
+ @infer.call(node.else_expr)
34
+ node.type = Builtins.common_type([node.then_expr.type, node.else_expr.type])
35
+ unless node.type
36
+ raise SignatureError,
37
+ "Conditional branches have incompatible types: #{node.then_expr.type} and #{node.else_expr.type}"
38
+ end
39
+ node
40
+ end
41
+
42
+ def infer_return(node)
43
+ @infer.call(node.expression) if node.expression
44
+ node.type = node.expression&.type
45
+ node
46
+ end
47
+
48
+ def infer_for_loop(node)
49
+ @infer.call(node.range_start)
50
+ @infer.call(node.range_end)
51
+ unless node.range_start.type == :int && node.range_end.type == :int
52
+ raise SignatureError,
53
+ "Loop bounds must be integers, got #{node.range_start.type.inspect} and #{node.range_end.type.inspect}"
54
+ end
55
+
56
+ @infer_in_scope.call(node.variable => :int) do
57
+ @infer.call(node.body)
58
+ end
59
+
60
+ node.type = nil
61
+ node
62
+ end
63
+
64
+ def infer_while_loop(node)
65
+ @infer.call(node.condition)
66
+ @infer.call(node.body, scoped: true)
67
+ node.type = nil
68
+ node
69
+ end
70
+
71
+ def infer_function_definition(node)
72
+ @infer_in_scope.call(node.param_types) do
73
+ @infer.call(node.body)
74
+
75
+ node.return_type ||= node.body&.type
76
+ validate_return_types!(node) if node.return_type
77
+ node.type = node.return_type
78
+ end
79
+
80
+ node
81
+ end
82
+
83
+ private
84
+
85
+ def validate_return_types!(function)
86
+ expressions = explicit_return_expressions(function.body)
87
+ expressions.concat(terminal_expressions(function.body))
88
+ expressions.uniq.each do |expression|
89
+ next if compatible_return?(function.return_type, expression)
90
+
91
+ raise SignatureError,
92
+ "Function #{function.name} returns #{return_type_of(expression).inspect}, expected #{function.return_type.inspect}"
93
+ end
94
+ end
95
+
96
+ def explicit_return_expressions(node)
97
+ case node
98
+ when IR::Block
99
+ node.statements.flat_map { |statement| explicit_return_expressions(statement) }
100
+ when IR::IfStatement
101
+ explicit_return_expressions(node.then_branch) + explicit_return_expressions(node.else_branch)
102
+ when IR::ForLoop, IR::WhileLoop
103
+ explicit_return_expressions(node.body)
104
+ when IR::Return
105
+ node.expression ? [node.expression] : []
106
+ else
107
+ []
108
+ end
109
+ end
110
+
111
+ def terminal_expressions(node)
112
+ case node
113
+ when IR::Block
114
+ terminal_expressions(node.statements.last)
115
+ when IR::IfStatement
116
+ terminal_expressions(node.then_branch) + terminal_expressions(node.else_branch)
117
+ when IR::Return, nil
118
+ []
119
+ else
120
+ node.type ? [node] : []
121
+ end
122
+ end
123
+
124
+ def compatible_return?(expected, expression)
125
+ if expected.is_a?(Array)
126
+ actual = return_type_of(expression)
127
+ return false unless actual.is_a?(Array) && actual.length == expected.length
128
+
129
+ return expected.zip(actual).all? { |expected_type, actual_type| compatible_type?(expected_type, actual_type) }
130
+ end
131
+
132
+ compatible_type?(expected, expression.type)
133
+ end
134
+
135
+ def return_type_of(expression)
136
+ return expression.elements.map(&:type) if expression.is_a?(IR::ArrayLiteral)
137
+ return expression.type.types if expression.type.is_a?(IR::TupleType)
138
+
139
+ expression.type
140
+ end
141
+
142
+ def compatible_type?(expected, actual)
143
+ expected == actual || (expected == :float && actual == :int)
144
+ end
145
+ end
146
+ end
147
+ end
@@ -0,0 +1,56 @@
1
+ # frozen_string_literal: true
2
+
3
+ module RLSL
4
+ module Prism
5
+ class DefinitionInferer
6
+ def initialize(infer:, lookup:, register:, collection_type_resolver:)
7
+ @infer = infer
8
+ @lookup = lookup
9
+ @register = register
10
+ @collection_type_resolver = collection_type_resolver
11
+ end
12
+
13
+ def infer_var_decl(node)
14
+ @infer.call(node.initializer) if node.initializer
15
+ node.type ||= node.initializer&.type
16
+ @register.call(node.name, node.type) if node.type
17
+ node
18
+ end
19
+
20
+ def infer_assignment(node)
21
+ @infer.call(node.target)
22
+ @infer.call(node.value)
23
+ existing_type = node.target.is_a?(IR::VarRef) ? @lookup.call(node.target.name) : node.target.type
24
+ if existing_type && node.value.type && !compatible_assignment?(existing_type, node.value.type)
25
+ raise SignatureError,
26
+ "Cannot assign #{node.value.type} to #{node.target.name} (#{existing_type})"
27
+ end
28
+
29
+ node.target.type ||= node.value.type
30
+ @register.call(node.target.name, node.value.type) if node.target.is_a?(IR::VarRef) && !existing_type
31
+ node.type = node.value.type
32
+ node
33
+ end
34
+
35
+ def infer_global_decl(node)
36
+ @infer.call(node.initializer) if node.initializer
37
+ node.type ||= @collection_type_resolver.resolve_global_decl(node)
38
+ @register.call(node.name, node.type) if node.type
39
+ node
40
+ end
41
+
42
+ def infer_multiple_assignment(node)
43
+ @infer.call(node.value)
44
+ @collection_type_resolver.assign_multiple_targets(node)
45
+ node.type = nil
46
+ node
47
+ end
48
+
49
+ private
50
+
51
+ def compatible_assignment?(target_type, value_type)
52
+ target_type == value_type || (target_type == :float && value_type == :int)
53
+ end
54
+ end
55
+ end
56
+ end