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
@@ -39,45 +39,102 @@ module RLSL
39
39
 
40
40
  def emit_var_decl(node)
41
41
  type = type_name(node.type || :float)
42
+ return "var #{node.name}: #{type}" unless node.initializer
43
+
44
+ if node.initializer.is_a?(IR::ArrayLiteral)
45
+ element_type_symbol = TypeShapes.element_type(node.type) || :float
46
+ element_type = type_name(element_type_symbol)
47
+ size = node.initializer.elements.length
48
+ values = node.initializer.elements.map do |element|
49
+ emit_typed_argument(element, element_type_symbol)
50
+ end.join(", ")
51
+ return "var #{node.name}: array<#{element_type}, #{size}> = array<#{element_type}, #{size}>(#{values})"
52
+ end
53
+
42
54
  value = emit(node.initializer)
43
- "let #{node.name}: #{type} = #{value}"
55
+ binding = node.mutable ? "var" : "let"
56
+ "#{binding} #{node.name}: #{type} = #{value}"
44
57
  end
45
58
 
46
59
  def emit_for_loop(node)
47
- var = node.variable
60
+ variable = node.variable
61
+ counter = loop_variable_mutated?(node) ? next_temporary_name("i") : variable
48
62
  start_val = emit(node.range_start)
49
63
  end_val = emit(node.range_end)
50
64
  body = emit_indented_block(node.body)
65
+ body = "#{indent} var #{variable}: i32 = #{counter};\n#{body}" if counter != variable
66
+
67
+ comparison = node.exclude_end ? "<" : "<="
68
+ return "for (var #{counter}: i32 = #{start_val}; #{counter} #{comparison} #{end_val}; #{counter}++) {\n#{body}#{indent}}" if node.range_end.is_a?(IR::Literal)
69
+
70
+ bound = next_temporary_name("end")
71
+ "let #{bound}: i32 = #{end_val};\n#{indent}for (var #{counter}: i32 = #{start_val}; #{counter} #{comparison} #{bound}; #{counter}++) {\n#{body}#{indent}}"
72
+ end
51
73
 
52
- "for (var #{var}: i32 = #{start_val}; #{var} < #{end_val}; #{var}++) {\n#{body}#{indent}}"
74
+ def emit_ternary(_node)
75
+ raise TargetCapabilityError,
76
+ "WGSL conditional expressions cannot be emitted without eager branch evaluation"
53
77
  end
54
78
 
55
- def emit_ternary(node)
56
- condition = emit(node.condition)
57
- then_expr = emit(node.then_expr)
58
- else_expr = emit(node.else_expr)
59
- "select(#{else_expr}, #{then_expr}, #{condition})"
79
+ def emit_binary_op(node)
80
+ if node.operator == "%" && node.type == :float
81
+ return "rlsl_mod(#{emit_float_operand(node.left)}, #{emit_float_operand(node.right)})"
82
+ end
83
+
84
+ super
85
+ end
86
+
87
+ def emit_func_call(node)
88
+ if node.name.to_sym == :mod
89
+ args = node.args.map { |argument| emit_float_operand(argument) }
90
+ return "rlsl_mod(#{args.join(', ')})"
91
+ end
92
+ if node.name.to_sym == :atan && node.args.length == 2
93
+ return emit_named_call("atan2", node.args, expected_types: node.expected_arg_types)
94
+ end
95
+
96
+ super
60
97
  end
61
98
 
62
99
  def emit_function_definition(node)
63
100
  name = node.name
64
- params = node.params.map do |param|
65
- param_type = type_name(node.param_types[param] || :float)
66
- "#{param}: #{param_type}"
101
+ mutable_params = mutated_parameters(node)
102
+ initializers = []
103
+ sampler_params = {}
104
+ params = node.params.flat_map do |param|
105
+ type = node.param_types[param] || :float
106
+ if type == :sampler2D
107
+ if mutable_params.include?(param)
108
+ raise TargetCapabilityError, "WGSL texture parameter #{param} cannot be reassigned"
109
+ end
110
+
111
+ sampler_params[param] = next_temporary_name("sampler")
112
+ ["#{param}: #{type_name(type)}", "#{sampler_params[param]}: sampler"]
113
+ elsif mutable_params.include?(param)
114
+ argument = next_temporary_name("param")
115
+ initializers << "#{indent} var #{param}: #{type_name(type)} = #{argument};\n"
116
+ ["#{argument}: #{type_name(type)}"]
117
+ else
118
+ ["#{param}: #{type_name(type)}"]
119
+ end
67
120
  end.join(", ")
68
121
 
69
122
  if node.return_type.is_a?(Array)
70
- @current_return_struct_name = "#{name}_result"
71
123
  struct_def = emit_result_struct(name, node.return_type)
72
- body = emit_indented_block(node.body, needs_return: true)
73
- @current_return_struct_name = nil
124
+ body = with_return_struct_name("#{name}_result") do
125
+ with_sampler_parameters(sampler_params) do
126
+ with_return_type(node.return_type) { emit_indented_block(node.body, needs_return: true) }
127
+ end
128
+ end
74
129
 
75
- "#{struct_def}fn #{name}(#{params}) -> #{name}_result {\n#{body}\n#{indent}}\n"
130
+ "#{struct_def}fn #{name}(#{params}) -> #{name}_result {\n#{initializers.join}#{body}#{indent}}\n"
76
131
  else
77
132
  return_type = type_name(node.return_type || :float)
78
- body = emit_indented_block(node.body, needs_return: true)
133
+ body = with_sampler_parameters(sampler_params) do
134
+ with_return_type(node.return_type) { emit_indented_block(node.body, needs_return: true) }
135
+ end
79
136
 
80
- "fn #{name}(#{params}) -> #{return_type} {\n#{body}\n#{indent}}\n"
137
+ "fn #{name}(#{params}) -> #{return_type} {\n#{initializers.join}#{body}#{indent}}\n"
81
138
  end
82
139
  end
83
140
 
@@ -85,7 +142,105 @@ module RLSL
85
142
  fields = types.each_with_index.map do |type, index|
86
143
  "#{indent}v#{index}: #{type_name(type)},"
87
144
  end.join("\n")
88
- "struct #{func_name}_result {\n#{fields}\n};\n"
145
+ "struct #{func_name}_result {\n#{fields}\n}\n"
146
+ end
147
+
148
+ def emit_hoisted_declarations(node)
149
+ node.hoisted_variables.map do |name, type|
150
+ "#{indent}var #{name}: #{type_name(type || :float)};\n"
151
+ end.join
152
+ end
153
+
154
+ def emit_global_decl(node)
155
+ if node.initializer.is_a?(IR::ArrayLiteral)
156
+ element_type_symbol = node.element_type || TypeShapes.element_type(node.initializer.type) || :float
157
+ element_type = type_name(element_type_symbol)
158
+ size = node.array_size || node.initializer.elements.length
159
+ values = node.initializer.elements.map do |element|
160
+ emit_typed_argument(element, element_type_symbol)
161
+ end.join(", ")
162
+ prefix = node.is_const ? "const" : "var<private>"
163
+ return "#{prefix} #{node.name}: array<#{element_type}, #{size}> = array<#{element_type}, #{size}>(#{values})"
164
+ end
165
+
166
+ prefix = node.is_const ? "const" : "var<private>"
167
+ "#{prefix} #{node.name}: #{type_name(node.type || :float)} = #{emit(node.initializer)}"
168
+ end
169
+
170
+ def emit_field_access(node)
171
+ return node.field.to_s if node.receiver.type == :uniforms && node.type == :sampler2D
172
+
173
+ code = super
174
+ return "(#{code} != 0)" if node.receiver.type == :uniforms && node.type == :bool
175
+
176
+ code
177
+ end
178
+
179
+ def emit_texture_call(name, node)
180
+ return unless profile.texture_functions.key?(name) && node.args.length >= 2
181
+
182
+ texture = emit(node.args[0])
183
+ sampler = sampler_name(node.args[0], texture)
184
+ uv = emit(node.args[1])
185
+ lod = node.args[2] ? emit_typed_argument(node.args[2], :float) : "0.0"
186
+ "textureSampleLevel(#{texture}, #{sampler}, #{uv}, #{lod})"
187
+ end
188
+
189
+ def emit_named_call(name, args, receiver: nil, expected_types: [])
190
+ arguments = receiver ? [receiver, *args] : args
191
+ rendered = arguments.each_with_index.flat_map do |argument, index|
192
+ if expected_types[index] == :sampler2D
193
+ texture = emit(argument)
194
+ [texture, sampler_name(argument, texture)]
195
+ else
196
+ [emit_typed_argument(argument, expected_types[index])]
197
+ end
198
+ end
199
+ "#{name}(#{rendered.join(', ')})"
200
+ end
201
+
202
+ def emit_multiple_assignment_target(target, declaration)
203
+ declaration ? "var #{target.name}: #{type_name(target.type || :float)}" : target.name.to_s
204
+ end
205
+
206
+ def emit_temporary_declaration(type, name, value)
207
+ "let #{name}: #{type} = #{value}"
208
+ end
209
+
210
+ def emit_tuple_value(node)
211
+ expected_types = Array(current_return_type)
212
+ elements = node.elements.each_with_index.map do |element, index|
213
+ emit_typed_argument(element, expected_types[index])
214
+ end.join(", ")
215
+ "#{current_return_struct_name}(#{elements})"
216
+ end
217
+
218
+ private
219
+
220
+ def mutated_parameters(node)
221
+ parameters = node.params.to_set
222
+ IR::Traversal.each(node.body).each_with_object(Set.new) do |current, names|
223
+ case current
224
+ when IR::Assignment
225
+ names.add(current.target.name) if current.target.is_a?(IR::VarRef) && parameters.include?(current.target.name)
226
+ when IR::MultipleAssignment
227
+ current.targets.each { |target| names.add(target.name) if parameters.include?(target.name) }
228
+ end
229
+ end
230
+ end
231
+
232
+ def with_sampler_parameters(parameters)
233
+ previous = @sampler_parameters
234
+ @sampler_parameters = parameters
235
+ yield
236
+ ensure
237
+ @sampler_parameters = previous
238
+ end
239
+
240
+ def sampler_name(argument, texture)
241
+ return @sampler_parameters[argument.name] if argument.is_a?(IR::VarRef) && @sampler_parameters&.key?(argument.name)
242
+
243
+ "#{texture}_sampler"
89
244
  end
90
245
  end
91
246
  end
@@ -0,0 +1,9 @@
1
+ # frozen_string_literal: true
2
+
3
+ require_relative "../errors"
4
+
5
+ module RLSL
6
+ module Prism
7
+ class UnsupportedSyntaxError < RLSL::Error; end
8
+ end
9
+ end
@@ -4,15 +4,16 @@ module RLSL
4
4
  module Prism
5
5
  module IR
6
6
  class IfStatement < Node
7
- attr_reader :condition, :then_branch, :else_branch
7
+ attr_reader :condition, :then_branch, :else_branch, :hoisted_variables
8
8
 
9
9
  visits :visit_if_statement
10
10
 
11
- def initialize(condition, then_branch, else_branch = nil, type = nil)
11
+ def initialize(condition, then_branch, else_branch = nil, type = nil, hoisted_variables: {})
12
12
  super()
13
13
  @condition = condition
14
14
  @then_branch = then_branch
15
15
  @else_branch = else_branch
16
+ @hoisted_variables = hoisted_variables
16
17
  @type = type
17
18
  end
18
19
  end
@@ -43,16 +44,17 @@ module RLSL
43
44
  end
44
45
 
45
46
  class ForLoop < Node
46
- attr_reader :variable, :range_start, :range_end, :body
47
+ attr_reader :variable, :range_start, :range_end, :body, :exclude_end
47
48
 
48
49
  visits :visit_for_loop
49
50
 
50
- def initialize(variable, range_start, range_end, body)
51
+ def initialize(variable, range_start, range_end, body, exclude_end: true)
51
52
  super()
52
53
  @variable = variable
53
54
  @range_start = range_start
54
55
  @range_end = range_end
55
56
  @body = body
57
+ @exclude_end = exclude_end
56
58
  @type = nil
57
59
  end
58
60
  end
@@ -39,14 +39,15 @@ module RLSL
39
39
  end
40
40
 
41
41
  class MultipleAssignment < Node
42
- attr_reader :targets, :value
42
+ attr_reader :targets, :value, :declarations
43
43
 
44
44
  visits :visit_multiple_assignment
45
45
 
46
- def initialize(targets, value)
46
+ def initialize(targets, value, declarations: nil)
47
47
  super()
48
48
  @targets = targets
49
49
  @value = value
50
+ @declarations = declarations || Array.new(targets.length, true)
50
51
  @type = nil
51
52
  end
52
53
  end
@@ -61,6 +62,20 @@ module RLSL
61
62
  def to_sym
62
63
  :"tuple_#{types.map(&:to_s).join('_')}"
63
64
  end
65
+
66
+ def ==(other)
67
+ other.is_a?(self.class) && types == other.types
68
+ end
69
+
70
+ alias eql? ==
71
+
72
+ def hash
73
+ [self.class, types].hash
74
+ end
75
+
76
+ def inspect
77
+ "#<#{self.class.name} #{types.inspect}>"
78
+ end
64
79
  end
65
80
  end
66
81
  end
@@ -16,14 +16,16 @@ module RLSL
16
16
 
17
17
  class VarDecl < Node
18
18
  attr_reader :name, :initializer
19
+ attr_accessor :mutable
19
20
 
20
21
  visits :visit_var_decl
21
22
 
22
- def initialize(name, initializer, type = nil)
23
+ def initialize(name, initializer, type = nil, mutable: false)
23
24
  super()
24
25
  @name = name
25
26
  @initializer = initializer
26
27
  @type = type
28
+ @mutable = mutable
27
29
  end
28
30
  end
29
31
 
@@ -92,6 +94,7 @@ module RLSL
92
94
 
93
95
  class FuncCall < Node
94
96
  attr_reader :name, :args, :receiver
97
+ attr_accessor :expected_arg_types
95
98
 
96
99
  visits :visit_func_call
97
100
 
@@ -100,6 +103,7 @@ module RLSL
100
103
  @name = name
101
104
  @args = args
102
105
  @receiver = receiver
106
+ @expected_arg_types = []
103
107
  @type = type
104
108
  end
105
109
  end
@@ -4,7 +4,7 @@ module RLSL
4
4
  module Prism
5
5
  module IR
6
6
  class Node
7
- attr_accessor :type
7
+ attr_accessor :type, :location
8
8
 
9
9
  def self.visits(method_name)
10
10
  define_method(:accept) do |visitor|
@@ -10,8 +10,12 @@ module RLSL
10
10
  return enum_for(:each, node) unless block_given?
11
11
  return if node.nil?
12
12
 
13
- yield node
14
- child_nodes(node).each { |child| each(child, &block) }
13
+ stack = [node]
14
+ until stack.empty?
15
+ current = stack.pop
16
+ yield current
17
+ stack.concat(child_nodes(current).reverse)
18
+ end
15
19
  end
16
20
 
17
21
  def child_nodes(node)
@@ -0,0 +1,30 @@
1
+ # frozen_string_literal: true
2
+
3
+ require "set"
4
+
5
+ require_relative "ir/traversal"
6
+
7
+ module RLSL
8
+ module Prism
9
+ class MutationAnalyzer
10
+ def analyze(node)
11
+ assigned_names = IR::Traversal.each(node).filter_map do |current|
12
+ case current
13
+ when IR::Assignment
14
+ current.target.name if current.target.is_a?(IR::VarRef)
15
+ when IR::MultipleAssignment
16
+ current.targets.select { |target| target.is_a?(IR::VarRef) }.map(&:name)
17
+ end
18
+ end.flatten.to_set
19
+
20
+ IR::Traversal.each(node) do |current|
21
+ next unless current.is_a?(IR::VarDecl)
22
+
23
+ current.mutable = assigned_names.include?(current.name)
24
+ end
25
+
26
+ node
27
+ end
28
+ end
29
+ end
30
+ end
@@ -0,0 +1,41 @@
1
+ # frozen_string_literal: true
2
+
3
+ module RLSL
4
+ module Prism
5
+ module NodeTraversal
6
+ module_function
7
+
8
+ def each(node)
9
+ return enum_for(__method__, node) unless block_given?
10
+ return unless node
11
+
12
+ stack = [node]
13
+ until stack.empty?
14
+ current = stack.pop
15
+ yield current
16
+ stack.concat(child_nodes(current).reverse)
17
+ end
18
+ end
19
+
20
+ def depth_exceeds?(node, maximum)
21
+ return false unless node
22
+
23
+ stack = [[node, 1]]
24
+ until stack.empty?
25
+ current, depth = stack.pop
26
+ return true if depth > maximum
27
+
28
+ child_nodes(current).reverse_each { |child| stack << [child, depth + 1] }
29
+ end
30
+ false
31
+ end
32
+
33
+ def child_nodes(node)
34
+ return node.compact_child_nodes if node.respond_to?(:compact_child_nodes)
35
+
36
+ Array(node.child_nodes).compact
37
+ end
38
+ private_class_method :child_nodes
39
+ end
40
+ end
41
+ end
@@ -0,0 +1,46 @@
1
+ # frozen_string_literal: true
2
+
3
+ require_relative "errors"
4
+
5
+ module RLSL
6
+ module Prism
7
+ module ParameterList
8
+ module_function
9
+
10
+ def required_names(parameter_container)
11
+ parameters = unwrap(parameter_container)
12
+ return [] unless parameters
13
+
14
+ unsupported = children(parameters).reject { |parameter| parameter.is_a?(::Prism::RequiredParameterNode) }
15
+ unless unsupported.empty?
16
+ raise UnsupportedSyntaxError, "Only required positional parameters are supported"
17
+ end
18
+
19
+ parameters.requireds.map(&:name)
20
+ end
21
+
22
+ def names(parameter_container)
23
+ parameters = unwrap(parameter_container)
24
+ return [] unless parameters
25
+
26
+ children(parameters).filter_map { |parameter| parameter.name if parameter.respond_to?(:name) }
27
+ end
28
+
29
+ def unwrap(container)
30
+ return container.parameters if container.respond_to?(:parameters)
31
+
32
+ container
33
+ end
34
+ private_class_method :unwrap
35
+
36
+ def children(parameters)
37
+ if parameters.respond_to?(:compact_child_nodes)
38
+ parameters.compact_child_nodes
39
+ else
40
+ Array(parameters.child_nodes).compact
41
+ end
42
+ end
43
+ private_class_method :children
44
+ end
45
+ end
46
+ end
@@ -0,0 +1,73 @@
1
+ # frozen_string_literal: true
2
+
3
+ require_relative "ir/traversal"
4
+
5
+ module RLSL
6
+ module Prism
7
+ class ReturnFlowError < RLSL::Error; end
8
+
9
+ class ReturnFlowValidator
10
+ VALUE_NODES = [
11
+ IR::VarDecl,
12
+ IR::Assignment,
13
+ IR::VarRef,
14
+ IR::Literal,
15
+ IR::BoolLiteral,
16
+ IR::BinaryOp,
17
+ IR::UnaryOp,
18
+ IR::FuncCall,
19
+ IR::FieldAccess,
20
+ IR::Swizzle,
21
+ IR::Ternary,
22
+ IR::Constant,
23
+ IR::Parenthesized,
24
+ IR::ArrayIndex
25
+ ].freeze
26
+
27
+ def validate!(node, needs_return:)
28
+ validate_functions!(node)
29
+ validate_returning_block!(node, "shader fragment") if needs_return
30
+ node
31
+ end
32
+
33
+ private
34
+
35
+ def validate_functions!(node)
36
+ IR::Traversal.each(node) do |current|
37
+ next unless current.is_a?(IR::FunctionDefinition)
38
+
39
+ validate_returning_block!(
40
+ current.body,
41
+ "function #{current.name}",
42
+ tuple_return: current.return_type.is_a?(Array)
43
+ )
44
+ end
45
+ end
46
+
47
+ def validate_returning_block!(node, context, tuple_return: false)
48
+ return if returns_value_on_all_paths?(node, tuple_return: tuple_return)
49
+
50
+ raise ReturnFlowError.new(
51
+ "#{context} does not return a value on every path"
52
+ ).with_source_location(node.location)
53
+ end
54
+
55
+ def returns_value_on_all_paths?(node, tuple_return: false)
56
+ case node
57
+ when IR::Block
58
+ returns_value_on_all_paths?(node.statements.last, tuple_return: tuple_return)
59
+ when IR::Return
60
+ !node.expression.nil? && (tuple_return || !node.expression.is_a?(IR::ArrayLiteral))
61
+ when IR::IfStatement
62
+ node.else_branch &&
63
+ returns_value_on_all_paths?(node.then_branch, tuple_return: tuple_return) &&
64
+ returns_value_on_all_paths?(node.else_branch, tuple_return: tuple_return)
65
+ when IR::ArrayLiteral
66
+ tuple_return
67
+ else
68
+ VALUE_NODES.any? { |klass| node.is_a?(klass) }
69
+ end
70
+ end
71
+ end
72
+ end
73
+ end
@@ -1,5 +1,8 @@
1
1
  # frozen_string_literal: true
2
2
 
3
+ require_relative "../node_traversal"
4
+ require_relative "../parameter_list"
5
+
3
6
  module RLSL
4
7
  module Prism
5
8
  class SourceExtractor
@@ -8,43 +11,38 @@ module RLSL
8
11
  extract_unit(source, start_line).to_source
9
12
  end
10
13
 
11
- def extract_unit(source, start_line)
14
+ def extract_unit(source, start_line, parameters: nil, source_name: "(shader block)")
12
15
  parsed = ::Prism.parse(source)
13
16
  raise SourceNotAvailable, "Unable to parse block source" unless parsed.success?
14
17
 
15
- block = block_at_line(parsed.value, start_line)
18
+ block = block_at_line(parsed.value, start_line, parameters)
16
19
  raise SourceNotAvailable, "Unable to locate block source" unless block
17
20
 
18
- SourceUnit.from_block(block)
21
+ SourceUnit.from_block(block, source_name: source_name)
19
22
  end
20
23
 
21
24
  private
22
25
 
23
- def block_at_line(node, start_line)
24
- each_node(node) do |current|
25
- next unless current.is_a?(::Prism::BlockNode)
26
- return current if current.location.start_line == start_line
26
+ def block_at_line(node, start_line, parameters)
27
+ candidates = NodeTraversal.each(node).select do |current|
28
+ current.is_a?(::Prism::BlockNode) && current.location.start_line == start_line
27
29
  end
30
+ candidates.select! { |candidate| parameter_names(candidate) == required_parameter_names(parameters) } if parameters
28
31
 
29
- nil
30
- end
31
- def each_node(node)
32
- return enum_for(:each_node, node) unless block_given?
33
- return unless node
34
-
35
- stack = [node]
36
-
37
- until stack.empty?
38
- current = stack.pop
39
- yield current
40
-
41
- children = if current.respond_to?(:compact_child_nodes)
42
- current.compact_child_nodes
43
- else
44
- Array(current.child_nodes).compact
45
- end
46
- stack.concat(children.reverse)
32
+ if candidates.length > 1
33
+ raise SourceNotAvailable,
34
+ "Multiple shader blocks start on line #{start_line}; put each block on its own line"
47
35
  end
36
+
37
+ candidates.first
38
+ end
39
+
40
+ def parameter_names(block)
41
+ ParameterList.names(block.parameters)
42
+ end
43
+
44
+ def required_parameter_names(parameters)
45
+ Array(parameters).filter_map { |_kind, name| name }
48
46
  end
49
47
  end
50
48
  end
@@ -8,7 +8,7 @@ require_relative "source_extractor/block_locator"
8
8
  module RLSL
9
9
  module Prism
10
10
  class SourceExtractor
11
- class SourceNotAvailable < StandardError; end
11
+ class SourceNotAvailable < RLSL::Error; end
12
12
 
13
13
  def initialize(block_locator = BlockLocator.new)
14
14
  @block_locator = block_locator
@@ -22,11 +22,12 @@ module RLSL
22
22
  file, line_num = block.source_location
23
23
  raise SourceNotAvailable, "Block source location not available" unless file && File.exist?(file)
24
24
 
25
- @block_locator.extract_unit(File.read(file), line_num)
26
- end
27
-
28
- def extract_from_string(source)
29
- source
25
+ @block_locator.extract_unit(
26
+ File.read(file),
27
+ line_num,
28
+ parameters: block.parameters,
29
+ source_name: file
30
+ )
30
31
  end
31
32
  end
32
33
  end