rlsl 1.0.0 → 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 (79) hide show
  1. checksums.yaml +4 -4
  2. data/.rubocop.yml +8 -0
  3. data/CHANGELOG.md +3 -1
  4. data/README.md +101 -28
  5. data/Rakefile +7 -0
  6. data/lib/rlsl/base_translator/call_parser.rb +13 -2
  7. data/lib/rlsl/base_translator/code_scanner.rb +39 -1
  8. data/lib/rlsl/base_translator.rb +24 -2
  9. data/lib/rlsl/code_generator/math_prelude.rb +103 -10
  10. data/lib/rlsl/code_generator/ruby_wrapper_generator.rb +52 -12
  11. data/lib/rlsl/code_generator/template_context.rb +6 -3
  12. data/lib/rlsl/code_generator/uniform_struct_generator.rb +1 -1
  13. data/lib/rlsl/code_generator.rb +7 -4
  14. data/lib/rlsl/compiled_shader.rb +1 -1
  15. data/lib/rlsl/errors.rb +29 -0
  16. data/lib/rlsl/glsl/translator.rb +19 -6
  17. data/lib/rlsl/msl/shader.rb +27 -9
  18. data/lib/rlsl/msl/translator.rb +25 -4
  19. data/lib/rlsl/msl/uniform_buffer_packer.rb +11 -7
  20. data/lib/rlsl/prism/ast_visitor/control_flow_visiting.rb +59 -18
  21. data/lib/rlsl/prism/ast_visitor/definition_visiting.rb +22 -9
  22. data/lib/rlsl/prism/ast_visitor/expression_visiting.rb +18 -14
  23. data/lib/rlsl/prism/ast_visitor/scope_context.rb +4 -0
  24. data/lib/rlsl/prism/ast_visitor.rb +78 -20
  25. data/lib/rlsl/prism/builtins/function_registry.rb +41 -24
  26. data/lib/rlsl/prism/builtins/operator_rules.rb +24 -0
  27. data/lib/rlsl/prism/builtins/swizzle_rules.rb +10 -1
  28. data/lib/rlsl/prism/builtins.rb +8 -0
  29. data/lib/rlsl/prism/emitters/base_emitter/control_flow_emission.rb +51 -8
  30. data/lib/rlsl/prism/emitters/base_emitter/definition_emission.rb +42 -12
  31. data/lib/rlsl/prism/emitters/base_emitter/expression_emission.rb +20 -4
  32. data/lib/rlsl/prism/emitters/base_emitter/statement_emission.rb +7 -1
  33. data/lib/rlsl/prism/emitters/base_emitter.rb +66 -5
  34. data/lib/rlsl/prism/emitters/c_emitter.rb +103 -7
  35. data/lib/rlsl/prism/emitters/glsl_emitter.rb +48 -0
  36. data/lib/rlsl/prism/emitters/msl_emitter.rb +43 -1
  37. data/lib/rlsl/prism/emitters/target_emitter.rb +56 -9
  38. data/lib/rlsl/prism/emitters/wgsl_emitter.rb +173 -18
  39. data/lib/rlsl/prism/errors.rb +9 -0
  40. data/lib/rlsl/prism/ir/control_flow.rb +6 -4
  41. data/lib/rlsl/prism/ir/definitions.rb +17 -2
  42. data/lib/rlsl/prism/ir/expressions.rb +5 -1
  43. data/lib/rlsl/prism/ir/node.rb +1 -1
  44. data/lib/rlsl/prism/ir/traversal.rb +6 -2
  45. data/lib/rlsl/prism/mutation_analyzer.rb +30 -0
  46. data/lib/rlsl/prism/node_traversal.rb +41 -0
  47. data/lib/rlsl/prism/parameter_list.rb +46 -0
  48. data/lib/rlsl/prism/return_flow_validator.rb +73 -0
  49. data/lib/rlsl/prism/source_extractor/block_locator.rb +23 -25
  50. data/lib/rlsl/prism/source_extractor.rb +7 -6
  51. data/lib/rlsl/prism/source_unit/parser.rb +42 -32
  52. data/lib/rlsl/prism/source_unit.rb +10 -10
  53. data/lib/rlsl/prism/target_capability_validator.rb +27 -2
  54. data/lib/rlsl/prism/transpiler.rb +40 -3
  55. data/lib/rlsl/prism/type_inference/call_type_resolver.rb +28 -4
  56. data/lib/rlsl/prism/type_inference/call_validator.rb +5 -1
  57. data/lib/rlsl/prism/type_inference/collection_type_resolver.rb +45 -19
  58. data/lib/rlsl/prism/type_inference/control_flow_inferer.rb +88 -7
  59. data/lib/rlsl/prism/type_inference/definition_inferer.rb +16 -1
  60. data/lib/rlsl/prism/type_inference/expression_inferer.rb +4 -0
  61. data/lib/rlsl/prism/type_inference/field_type_resolver.rb +23 -2
  62. data/lib/rlsl/prism/type_inference/scope_stack.rb +1 -1
  63. data/lib/rlsl/prism/type_inference/type_shapes.rb +3 -3
  64. data/lib/rlsl/prism/type_inference.rb +10 -5
  65. data/lib/rlsl/shader_builder/build_service.rb +34 -7
  66. data/lib/rlsl/shader_builder/native_extension_compiler.rb +50 -24
  67. data/lib/rlsl/shader_builder/shader_definition.rb +3 -3
  68. data/lib/rlsl/shader_builder/source_resolver.rb +19 -1
  69. data/lib/rlsl/shader_builder.rb +50 -10
  70. data/lib/rlsl/shader_name.rb +18 -0
  71. data/lib/rlsl/types/type_spec.rb +3 -3
  72. data/lib/rlsl/types/value_normalizer.rb +16 -11
  73. data/lib/rlsl/types.rb +0 -4
  74. data/lib/rlsl/uniform_context.rb +18 -1
  75. data/lib/rlsl/version.rb +1 -1
  76. data/lib/rlsl/wgsl/translator.rb +29 -6
  77. data/lib/rlsl/wgsl/uniform_layout.rb +25 -0
  78. data/lib/rlsl.rb +24 -3
  79. metadata +23 -11
@@ -2,23 +2,33 @@
2
2
 
3
3
  require "prism"
4
4
 
5
+ require_relative "../node_traversal"
6
+ require_relative "../parameter_list"
7
+
5
8
  module RLSL
6
9
  module Prism
7
10
  class SourceUnitParser
8
- def initialize(source)
11
+ def initialize(source, source_name: "(shader source)")
9
12
  @source = source.to_s
13
+ @source_name = source_name
10
14
  end
11
15
 
12
16
  def parse
13
- normalized = @source.strip
14
- return SourceUnit.new(params: [], body: "") if normalized.empty?
17
+ normalized, leading_line_offset = strip_with_line_offset(@source)
18
+ if normalized.empty?
19
+ return SourceUnit.new(params: [], body: "", source_name: @source_name, line_offset: leading_line_offset)
20
+ end
15
21
 
16
- params_source, body_source = split_sections(normalized)
17
- validate_body!(body_source)
22
+ params_source, body_source, parameter_line_offset = split_sections(normalized)
23
+ stripped_body, body_line_offset = strip_with_line_offset(body_source)
24
+ line_offset = leading_line_offset + parameter_line_offset + body_line_offset
25
+ validate_body!(stripped_body, line_offset)
18
26
 
19
27
  SourceUnit.new(
20
28
  params: parse_params(params_source),
21
- body: body_source.strip
29
+ body: stripped_body,
30
+ source_name: @source_name,
31
+ line_offset: line_offset
22
32
  )
23
33
  end
24
34
 
@@ -27,51 +37,51 @@ module RLSL
27
37
  def split_sections(source)
28
38
  lines = source.lines
29
39
  first_line = lines.first&.strip
30
- return [nil, source] unless parameter_line?(first_line)
40
+ return [nil, source, 0] unless parameter_line?(first_line)
31
41
 
32
- [first_line, lines[1..].to_a.join]
42
+ [first_line, lines[1..].to_a.join, 1]
33
43
  end
34
44
 
35
45
  def parameter_line?(line)
36
- line&.start_with?("|") && line&.end_with?("|")
46
+ line&.start_with?("|") && line.end_with?("|")
37
47
  end
38
48
 
39
- def validate_body!(body_source)
49
+ def validate_body!(body_source, line_offset)
40
50
  return if body_source.to_s.strip.empty?
41
51
 
42
52
  parsed = ::Prism.parse(body_source)
43
- raise ArgumentError, "Unable to parse source unit body" unless parsed.success?
53
+ return if parsed.success?
54
+
55
+ error = RLSL::ParseError.new("Unable to parse source unit body")
56
+ location = parsed.errors.first&.location
57
+ if location
58
+ error.with_source_location(
59
+ RLSL::SourceLocation.new(
60
+ source_name: @source_name,
61
+ line: line_offset + location.start_line,
62
+ column: location.start_column + 1
63
+ )
64
+ )
65
+ end
66
+ raise error
44
67
  end
45
68
 
46
69
  def parse_params(params_source)
47
70
  return [] unless params_source
48
71
 
49
72
  parsed = ::Prism.parse("proc do #{params_source}\nend\n")
50
- raise ArgumentError, "Unable to parse source unit params" unless parsed.success?
73
+ raise RLSL::ParseError, "Unable to parse source unit params" unless parsed.success?
51
74
 
52
- block = each_node(parsed.value).find { |node| node.is_a?(::Prism::BlockNode) }
53
- return [] unless block&.parameters
75
+ block = NodeTraversal.each(parsed.value).find { |node| node.is_a?(::Prism::BlockNode) }
76
+ return [] unless block
54
77
 
55
- block.parameters.parameters.requireds.map(&:name)
78
+ ParameterList.required_names(block.parameters)
56
79
  end
57
80
 
58
- def each_node(node)
59
- return enum_for(:each_node, node) unless block_given?
60
- return unless node
61
-
62
- stack = [node]
63
-
64
- until stack.empty?
65
- current = stack.pop
66
- yield current
67
-
68
- children = if current.respond_to?(:compact_child_nodes)
69
- current.compact_child_nodes
70
- else
71
- Array(current.child_nodes).compact
72
- end
73
- stack.concat(children.reverse)
74
- end
81
+ def strip_with_line_offset(source)
82
+ text = source.to_s
83
+ leading_whitespace = text[/\A\s*/].to_s
84
+ [text.strip, leading_whitespace.count("\n")]
75
85
  end
76
86
  end
77
87
  end
@@ -1,34 +1,34 @@
1
1
  # frozen_string_literal: true
2
2
 
3
3
  require_relative "source_unit/parser"
4
+ require_relative "parameter_list"
4
5
 
5
6
  module RLSL
6
7
  module Prism
7
- SourceUnit = Struct.new(:params, :body, keyword_init: true) do
8
+ SourceUnit = Struct.new(:params, :body, :source_name, :line_offset, keyword_init: true) do
8
9
  class << self
9
- def from_source(source)
10
- SourceUnitParser.new(source).parse
10
+ def from_source(source, source_name: "(shader source)")
11
+ SourceUnitParser.new(source, source_name: source_name).parse
11
12
  end
12
13
 
13
- def from_block(block)
14
+ def from_block(block, source_name: "(shader block)")
14
15
  new(
15
16
  params: extract_params(block),
16
- body: block.body&.slice.to_s.strip
17
+ body: block.body&.slice.to_s.strip,
18
+ source_name: source_name,
19
+ line_offset: block.body ? block.body.location.start_line - 1 : block.location.start_line - 1
17
20
  )
18
21
  end
19
22
 
20
23
  private
21
24
 
22
25
  def extract_params(block)
23
- return [] unless block.parameters
24
-
25
- parameters = block.parameters.parameters
26
- parameters.requireds.map(&:name)
26
+ ParameterList.required_names(block.parameters)
27
27
  end
28
28
  end
29
29
 
30
30
  def without_params
31
- self.class.new(params: [], body: body)
31
+ self.class.new(params: [], body: body, source_name: source_name, line_offset: line_offset)
32
32
  end
33
33
 
34
34
  def to_source
@@ -6,14 +6,19 @@ require_relative "type_inference/type_shapes"
6
6
 
7
7
  module RLSL
8
8
  module Prism
9
- class TargetCapabilityError < StandardError; end
9
+ class TargetCapabilityError < RLSL::Error; end
10
10
 
11
11
  class TargetCapabilityValidator
12
12
  include TypeShapes
13
13
 
14
14
  def validate!(node, target)
15
15
  @target = target.to_sym
16
- IR::Traversal.each(node) { |current| validate_node!(current) }
16
+ IR::Traversal.each(node) do |current|
17
+ validate_node!(current)
18
+ rescue RLSL::Error => error
19
+ error.with_source_location(current.location)
20
+ raise
21
+ end
17
22
  node
18
23
  end
19
24
 
@@ -49,6 +54,26 @@ module RLSL
49
54
  Builtins.explicit_types(node.name).each do |type|
50
55
  validate_type!(type, context: "builtin #{node.name}")
51
56
  end
57
+
58
+ validate_c_builtin_overload!(node) if target == :c
59
+ end
60
+
61
+ def validate_c_builtin_overload!(node)
62
+ vector_types = node.args.map(&:type).select { |type| Builtins.vector_type?(type) }
63
+ scalar_only = %i[sqrt abs sign floor ceil fract mod min max clamp step smoothstep]
64
+ if scalar_only.include?(node.name.to_sym) && !vector_types.empty?
65
+ raise TargetCapabilityError,
66
+ "Builtin #{node.name} does not support vector arguments on C"
67
+ end
68
+
69
+ if %i[length normalize dot distance].include?(node.name.to_sym) && vector_types.empty?
70
+ raise TargetCapabilityError, "Builtin #{node.name} requires vector arguments on C"
71
+ end
72
+
73
+ return unless %i[reflect refract].include?(node.name.to_sym)
74
+ return if node.args.first&.type == :vec3 && node.args[1]&.type == :vec3
75
+
76
+ raise TargetCapabilityError, "Builtin #{node.name} currently requires vec3 arguments on C"
52
77
  end
53
78
 
54
79
  def validate_type!(type, context:)
@@ -10,6 +10,8 @@ require_relative "builtins"
10
10
  require_relative "ast_visitor"
11
11
  require_relative "type_inference"
12
12
  require_relative "target_capability_validator"
13
+ require_relative "mutation_analyzer"
14
+ require_relative "return_flow_validator"
13
15
  require_relative "emitters/base_emitter"
14
16
  require_relative "emitters/target_emitter"
15
17
  require_relative "emitters/c_emitter"
@@ -51,9 +53,17 @@ module RLSL
51
53
  )
52
54
  end
53
55
 
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
62
+
54
63
  def emit(target, compilation:, needs_return: true)
55
64
  emitter = resolve_emitter(target)
56
65
  validate_target_capabilities!(compilation.ir, target)
66
+ ReturnFlowValidator.new.validate!(compilation.ir, needs_return: needs_return)
57
67
  emitter.emit(compilation.ir, needs_return: needs_return)
58
68
  end
59
69
 
@@ -73,20 +83,36 @@ module RLSL
73
83
  )
74
84
  end
75
85
 
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
93
+
76
94
  private
77
95
 
78
96
  def source_unit(source)
97
+ return source if source.is_a?(SourceUnit)
98
+
79
99
  SourceUnit.from_source(source)
80
100
  end
81
101
 
82
102
  def build_ir(unit)
83
- visitor = ASTVisitor.new(uniforms: @uniforms, params: unit.params)
103
+ visitor = ASTVisitor.new(
104
+ uniforms: @uniforms,
105
+ params: unit.params,
106
+ source_name: unit.source_name,
107
+ line_offset: unit.line_offset
108
+ )
84
109
  visitor.parse(unit.body)
85
110
  end
86
111
 
87
112
  def compile_unit(unit, function_signatures: nil)
88
113
  ir = build_ir(unit)
89
114
  apply_function_signatures(ir, function_signatures || {})
115
+ MutationAnalyzer.new.analyze(ir)
90
116
  infer_ir(ir)
91
117
  CompilationUnit.new(source_unit: unit, ir: ir)
92
118
  end
@@ -104,7 +130,7 @@ module RLSL
104
130
 
105
131
  def resolve_emitter(target)
106
132
  emitter_class = TARGETS[target.to_sym]
107
- raise "Unknown target: #{target}" unless emitter_class
133
+ raise RLSL::Error, "Unknown target: #{target}" unless emitter_class
108
134
 
109
135
  emitter_class.new
110
136
  end
@@ -120,10 +146,21 @@ module RLSL
120
146
  next unless stmt.is_a?(IR::FunctionDefinition)
121
147
 
122
148
  sig = signatures[stmt.name]
123
- 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
124
154
 
125
155
  stmt.return_type = sig[:returns]
126
156
  stmt.param_types = sig[:params] || {}
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)
163
+ end
127
164
  end
128
165
  end
129
166
  end
@@ -15,21 +15,45 @@ module RLSL
15
15
  custom_function = @custom_functions[node.name.to_sym]
16
16
  return resolve_custom(node, custom_function) if custom_function
17
17
 
18
- node.receiver&.type
18
+ raise SignatureError, "Unknown shader function #{node.name}"
19
19
  end
20
20
 
21
21
  private
22
22
 
23
23
  def resolve_builtin(node, signature)
24
- arg_types = node.args.map(&:type)
24
+ arg_types = effective_arg_types(node)
25
25
  @call_validator.validate_builtin!(node, arg_types, signature)
26
- Builtins.resolve_return_type(signature[:returns], arg_types)
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
27
33
  end
28
34
 
29
35
  def resolve_custom(node, signature)
30
- @call_validator.validate_custom!(node.name, node.args.map(&:type), signature)
36
+ @call_validator.validate_custom!(node.name, effective_arg_types(node), signature)
37
+ node.expected_arg_types = signature[:params]&.values || []
31
38
  signature[:returns]
32
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
33
57
  end
34
58
  end
35
59
  end
@@ -15,7 +15,11 @@ module RLSL
15
15
 
16
16
  def validate_custom!(name, arg_types, signature)
17
17
  params = signature[:params]
18
- return unless params
18
+ unless params
19
+ return if arg_types.empty?
20
+
21
+ raise SignatureError, "Function #{name} requires explicit parameter types"
22
+ end
19
23
 
20
24
  validate_signature!(name, arg_types, params.values)
21
25
  end
@@ -12,8 +12,13 @@ module RLSL
12
12
  end
13
13
 
14
14
  def resolve_array_literal(node)
15
- element_type = node.elements.first&.type || :float
16
- TypeShapes.array(element_type)
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)
17
22
  end
18
23
 
19
24
  def resolve_array_index(node)
@@ -27,9 +32,8 @@ module RLSL
27
32
  def resolve_global_decl(node)
28
33
  if node.initializer.is_a?(IR::ArrayLiteral)
29
34
  node.array_size ||= node.initializer.elements.length
30
- first_elem = node.initializer.elements.first
31
- node.element_type ||= first_elem&.type || :float
32
- return TypeShapes.array(node.element_type)
35
+ node.element_type ||= TypeShapes.element_type(node.initializer.type)
36
+ return TypeShapes.array(node.element_type, node.array_size)
33
37
  end
34
38
 
35
39
  node.initializer&.type
@@ -37,41 +41,63 @@ module RLSL
37
41
 
38
42
  def assign_multiple_targets(node)
39
43
  value_type = node.value.type
40
- return assign_tuple_targets(node.targets, value_type.types) if value_type.is_a?(IR::TupleType)
41
- return assign_array_targets(node.targets, TypeShapes.element_type(value_type)) if TypeShapes.array?(value_type)
42
- return assign_custom_targets(node.targets, @custom_functions[node.value.name]) if custom_multi_return?(node.value)
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)
43
48
 
44
49
  nil
45
50
  end
46
51
 
47
52
  private
48
53
 
49
- def assign_tuple_targets(targets, types)
50
- targets.each_with_index do |target, index|
51
- assign_target(target, types[index])
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])
52
58
  end
53
59
  end
54
60
 
55
- def assign_array_targets(targets, type)
56
- targets.each do |target|
57
- assign_target(target, type)
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)
58
66
  end
59
67
  end
60
68
 
61
- def assign_custom_targets(targets, signature)
69
+ def assign_custom_targets(node, signature)
62
70
  returns = signature[:returns]
63
71
  return unless returns.is_a?(Array)
64
72
 
65
- targets.each_with_index do |target, index|
66
- assign_target(target, returns[index])
67
- end
73
+ assign_tuple_targets(node, returns)
68
74
  end
69
75
 
70
- def assign_target(target, type)
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
+
71
87
  target.type = type
72
88
  @register.call(target.name, type)
73
89
  end
74
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
+
75
101
  def custom_multi_return?(value)
76
102
  value.is_a?(IR::FuncCall) && @custom_functions.key?(value.name)
77
103
  end
@@ -3,18 +3,27 @@
3
3
  module RLSL
4
4
  module Prism
5
5
  class ControlFlowInferer
6
- def initialize(infer:, infer_child_scope:, infer_in_scope:)
6
+ def initialize(infer:, infer_in_scope:, lookup:)
7
7
  @infer = infer
8
- @infer_child_scope = infer_child_scope
9
8
  @infer_in_scope = infer_in_scope
9
+ @lookup = lookup
10
10
  end
11
11
 
12
12
  def infer_if_statement(node)
13
13
  @infer.call(node.condition)
14
- @infer_child_scope.call(node.then_branch)
15
- @infer_child_scope.call(node.else_branch) if node.else_branch
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
16
25
 
17
- node.type = node.then_branch.type
26
+ node.type = node.then_branch&.type
18
27
  node
19
28
  end
20
29
 
@@ -22,7 +31,11 @@ module RLSL
22
31
  @infer.call(node.condition)
23
32
  @infer.call(node.then_expr)
24
33
  @infer.call(node.else_expr)
25
- node.type = node.then_expr.type
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
26
39
  node
27
40
  end
28
41
 
@@ -35,6 +48,10 @@ module RLSL
35
48
  def infer_for_loop(node)
36
49
  @infer.call(node.range_start)
37
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
38
55
 
39
56
  @infer_in_scope.call(node.variable => :int) do
40
57
  @infer.call(node.body)
@@ -46,7 +63,7 @@ module RLSL
46
63
 
47
64
  def infer_while_loop(node)
48
65
  @infer.call(node.condition)
49
- @infer_child_scope.call(node.body)
66
+ @infer.call(node.body, scoped: true)
50
67
  node.type = nil
51
68
  node
52
69
  end
@@ -56,11 +73,75 @@ module RLSL
56
73
  @infer.call(node.body)
57
74
 
58
75
  node.return_type ||= node.body&.type
76
+ validate_return_types!(node) if node.return_type
59
77
  node.type = node.return_type
60
78
  end
61
79
 
62
80
  node
63
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
64
145
  end
65
146
  end
66
147
  end
@@ -3,8 +3,9 @@
3
3
  module RLSL
4
4
  module Prism
5
5
  class DefinitionInferer
6
- def initialize(infer:, register:, collection_type_resolver:)
6
+ def initialize(infer:, lookup:, register:, collection_type_resolver:)
7
7
  @infer = infer
8
+ @lookup = lookup
8
9
  @register = register
9
10
  @collection_type_resolver = collection_type_resolver
10
11
  end
@@ -19,6 +20,14 @@ module RLSL
19
20
  def infer_assignment(node)
20
21
  @infer.call(node.target)
21
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
22
31
  node.type = node.value.type
23
32
  node
24
33
  end
@@ -36,6 +45,12 @@ module RLSL
36
45
  node.type = nil
37
46
  node
38
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
39
54
  end
40
55
  end
41
56
  end
@@ -65,6 +65,10 @@ module RLSL
65
65
 
66
66
  def infer_swizzle(node)
67
67
  @infer.call(node.receiver)
68
+ unless Builtins.valid_swizzle_for_type?(node.components, node.receiver.type)
69
+ raise SignatureError, "Invalid swizzle #{node.components.inspect} for #{node.receiver.type || :unknown}"
70
+ end
71
+
68
72
  node.type = Builtins.swizzle_type(node.components)
69
73
  node
70
74
  end