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,33 +3,33 @@
3
3
  module RLSL
4
4
  module MSL
5
5
  class Translator < BaseTranslator
6
- TYPE_MAP = {
7
- "vec2" => "float2",
8
- "vec3" => "float3",
9
- "vec4" => "float4"
10
- }.freeze
11
-
12
- FUNC_REPLACEMENTS = BaseTranslator.common_func_replacements(
13
- target_vec2: "float2",
14
- target_vec3: "float3",
15
- target_vec4: "float4"
16
- ).freeze
6
+ PROFILE = BaseTranslator.build_profile(
7
+ uniform_target: :msl,
8
+ identifier_replacements: {
9
+ "vec2" => "float2",
10
+ "vec3" => "float3",
11
+ "vec4" => "float4"
12
+ },
13
+ call_rewrites: BaseTranslator.common_call_rewrites(
14
+ target_vec2: "float2",
15
+ target_vec3: "float3",
16
+ target_vec4: "float4"
17
+ )
18
+ )
17
19
 
18
20
  protected
19
21
 
20
- def translate_code(c_code)
21
- result = super(c_code)
22
- return result if result.empty?
23
-
24
- result.gsub!(/\bstatic\s+/, "")
25
- result.gsub!(/\binline\s+/, "")
26
- result
27
- end
28
-
29
22
  def generate_shader(helpers, fragment)
30
23
  <<~MSL
31
24
  #include <metal_stdlib>
32
25
  using namespace metal;
26
+ #{generated_by_comment("MSL")}
27
+
28
+ constexpr sampler rlsl_texture_sampler(coord::normalized, address::clamp_to_edge, filter::linear);
29
+
30
+ float rlsl_mod(float x, float y) {
31
+ return x - y * floor(x / y);
32
+ }
33
33
 
34
34
  // Uniform buffer structure
35
35
  struct Uniforms {
@@ -40,22 +40,22 @@ module RLSL
40
40
  #{helpers}
41
41
 
42
42
  // Fragment shader function
43
- float3 shader_fragment(float2 frag_coord, float2 resolution, constant Uniforms& u) {
44
- float2 uv = frag_coord / resolution.y;
45
- #{fragment}
43
+ float3 shader_fragment(float2 frag_coord, float2 resolution, constant Uniforms& u#{fragment_texture_parameters}) {
44
+ #{indent_source(fragment, 4)}
46
45
  }
47
46
 
48
47
  // Compute kernel entry point
49
48
  kernel void compute_shader(
50
49
  texture2d<float, access::write> output [[texture(0)]],
51
50
  constant Uniforms& u [[buffer(0)]],
51
+ #{kernel_texture_parameters}
52
52
  uint2 gid [[thread_position_in_grid]]
53
53
  ) {
54
54
  float2 resolution = float2(output.get_width(), output.get_height());
55
55
  if (gid.x >= uint(resolution.x) || gid.y >= uint(resolution.y)) return;
56
56
 
57
57
  float2 frag_coord = float2(gid.x, resolution.y - 1.0 - float(gid.y));
58
- float3 color = shader_fragment(frag_coord, resolution, u);
58
+ float3 color = shader_fragment(frag_coord, resolution, u#{fragment_texture_arguments});
59
59
 
60
60
  output.write(float4(clamp(color, 0.0, 1.0), 1.0), gid);
61
61
  }
@@ -64,25 +64,29 @@ module RLSL
64
64
 
65
65
  private
66
66
 
67
+ def profile
68
+ PROFILE
69
+ end
70
+
67
71
  def generate_uniform_struct
68
- fields = ["float2 resolution;"]
69
- @uniforms.each do |name, type|
70
- msl_type = uniform_type_to_target(type)
71
- fields << "#{msl_type} #{name};"
72
+ fields = uniform_lines(resolution_line: "#{target_vec2_type} resolution;") do |name, msl_type|
73
+ "#{msl_type} #{name};"
72
74
  end
73
75
  fields.join("\n ")
74
76
  end
75
77
 
76
- def target_vec2_type
77
- "float2"
78
+ def fragment_texture_parameters
79
+ texture_uniforms.keys.map { |name| ", texture2d<float> #{name}" }.join
78
80
  end
79
81
 
80
- def target_vec3_type
81
- "float3"
82
+ def fragment_texture_arguments
83
+ texture_uniforms.keys.map { |name| ", #{name}" }.join
82
84
  end
83
85
 
84
- def target_vec4_type
85
- "float4"
86
+ def kernel_texture_parameters
87
+ texture_uniforms.keys.each_with_index.map do |name, index|
88
+ " texture2d<float, access::sample> #{name} [[texture(#{index + 1})]],"
89
+ end.join("\n")
86
90
  end
87
91
  end
88
92
  end
@@ -0,0 +1,72 @@
1
+ # frozen_string_literal: true
2
+
3
+ module RLSL
4
+ module MSL
5
+ class UniformBufferPacker
6
+ def initialize(shader_name, uniform_types, uniform_names)
7
+ @shader_name = shader_name
8
+ @uniform_types = uniform_types
9
+ @uniform_names = uniform_names
10
+ end
11
+
12
+ def pack(width, height, uniforms)
13
+ normalized_uniforms = UniformTypes.normalize_values(@uniform_types, uniforms, shader_name: @shader_name)
14
+ data = [width.to_f, height.to_f].pack("e2")
15
+ current_offset = 8
16
+
17
+ @uniform_names.each do |name|
18
+ value = normalized_uniforms[name]
19
+ spec = UniformTypes.metal_spec(@uniform_types[name])
20
+ current_offset, data = append_padding(data, current_offset, spec.metal_alignment)
21
+ data << pack_uniform_value(spec, value)
22
+ current_offset += spec.metal_size
23
+ end
24
+
25
+ if data.bytesize > 256
26
+ raise ArgumentError, "Metal uniform buffer exceeds 256 bytes (#{data.bytesize} bytes)"
27
+ end
28
+
29
+ data.ljust(256, "\x00")
30
+ end
31
+
32
+ private
33
+
34
+ def append_padding(data, current_offset, alignment)
35
+ padding_needed = (alignment - (current_offset % alignment)) % alignment
36
+ return [current_offset, data] if padding_needed.zero?
37
+
38
+ [current_offset + padding_needed, data + ("\x00" * padding_needed)]
39
+ end
40
+
41
+ def pack_uniform_value(spec, value)
42
+ case spec.wrapper_kind
43
+ when :float
44
+ [value.to_f].pack("e")
45
+ when :int
46
+ [value.to_i].pack("l<")
47
+ when :bool
48
+ [value ? 1 : 0].pack("l<")
49
+ when :vector
50
+ pack_vector_uniform(spec.vector_size, value)
51
+ else
52
+ raise ArgumentError, "Unsupported Metal uniform type: #{spec.c_type}"
53
+ end
54
+ end
55
+
56
+ def pack_vector_uniform(vector_size, value)
57
+ components = Array(value).map(&:to_f)
58
+
59
+ case vector_size
60
+ when 2
61
+ components.pack("e2")
62
+ when 3
63
+ (components + [0.0]).pack("e4")
64
+ when 4
65
+ components.pack("e4")
66
+ else
67
+ raise ArgumentError, "Unsupported vector size: #{vector_size}"
68
+ end
69
+ end
70
+ end
71
+ end
72
+ end
@@ -0,0 +1,137 @@
1
+ # frozen_string_literal: true
2
+
3
+ module RLSL
4
+ module Prism
5
+ class ASTVisitor
6
+ module ControlFlowVisiting
7
+ VISITORS = {}.tap do |visitors|
8
+ visitors[::Prism::BlockNode] = :visit_block if defined?(::Prism::BlockNode)
9
+ visitors[::Prism::LambdaNode] = :visit_lambda if defined?(::Prism::LambdaNode)
10
+ visitors[::Prism::IfNode] = :visit_if if defined?(::Prism::IfNode)
11
+ visitors[::Prism::ElseNode] = :visit_else if defined?(::Prism::ElseNode)
12
+ visitors[::Prism::UnlessNode] = :visit_unless if defined?(::Prism::UnlessNode)
13
+ visitors[::Prism::ReturnNode] = :visit_return if defined?(::Prism::ReturnNode)
14
+ visitors[::Prism::RangeNode] = :visit_range if defined?(::Prism::RangeNode)
15
+ visitors[::Prism::ForNode] = :visit_for if defined?(::Prism::ForNode)
16
+ visitors[::Prism::WhileNode] = :visit_while if defined?(::Prism::WhileNode)
17
+ visitors[::Prism::BreakNode] = :visit_break if defined?(::Prism::BreakNode)
18
+ end.freeze
19
+
20
+ private
21
+
22
+ def visit_block(node)
23
+ visit_with_scoped_vars(node.body, params: extract_block_params(node))
24
+ end
25
+
26
+ def visit_lambda(node)
27
+ visit_block(node)
28
+ end
29
+
30
+ def visit_if(node)
31
+ condition = visit(node.predicate)
32
+ hoisted_variables = hoist_branch_variables(node)
33
+ then_branch = visit_with_scoped_vars(node.statements)
34
+ else_branch = node.subsequent ? visit_with_scoped_vars(node.subsequent) : nil
35
+
36
+ IR::IfStatement.new(condition, then_branch, else_branch, hoisted_variables: hoisted_variables)
37
+ end
38
+
39
+ def visit_else(node)
40
+ visit(node.statements)
41
+ end
42
+
43
+ def visit_unless(node)
44
+ condition = IR::UnaryOp.new("!", visit(node.predicate))
45
+ hoisted_variables = hoist_branch_variables(node)
46
+ then_branch = visit_with_scoped_vars(node.statements)
47
+ else_branch = node.else_clause ? visit_with_scoped_vars(node.else_clause) : nil
48
+
49
+ IR::IfStatement.new(condition, then_branch, else_branch, hoisted_variables: hoisted_variables)
50
+ end
51
+
52
+ def visit_return(node)
53
+ arguments = node.arguments&.arguments || []
54
+ if arguments.length > 1
55
+ raise UnsupportedSyntaxError, "Returning multiple values requires an explicitly declared tuple helper"
56
+ end
57
+
58
+ expr = arguments.empty? ? nil : normalize_expression(visit(arguments.first))
59
+ IR::Return.new(expr)
60
+ end
61
+
62
+ def visit_range(node)
63
+ [visit(node.left), visit(node.right), node.exclude_end?]
64
+ end
65
+
66
+ def visit_for(node)
67
+ range = visit(node.collection)
68
+ variable = node.index.name.to_sym
69
+ body = visit_with_scoped_vars(node.statements, params: [variable]) || IR::Block.new
70
+ IR::ForLoop.new(variable, range[0], range[1], body, exclude_end: range[2])
71
+ end
72
+
73
+ def visit_call_with_block(node)
74
+ unless times_loop?(node)
75
+ raise UnsupportedSyntaxError, "Blocks are only supported for Integer#times loops"
76
+ end
77
+
78
+ count = visit(node.receiver)
79
+ block_params = extract_block_params(node.block)
80
+ var_name = block_params.first || next_implicit_loop_variable
81
+ block = visit_with_scoped_vars(node.block.body, params: [var_name]) || IR::Block.new
82
+ IR::ForLoop.new(var_name, IR::Literal.new(0, :int), count, block)
83
+ end
84
+
85
+ def visit_while(node)
86
+ condition = visit(node.predicate)
87
+ body = visit_with_scoped_vars(node.statements) || IR::Block.new
88
+ IR::WhileLoop.new(condition, body)
89
+ end
90
+
91
+ def visit_break(_node)
92
+ IR::Break.new
93
+ end
94
+
95
+ def times_loop?(node)
96
+ node.name.to_s == "times" && node.receiver
97
+ end
98
+
99
+ def hoist_branch_variables(node)
100
+ then_names = definitely_assigned_names(node.statements)
101
+ else_node = node.respond_to?(:subsequent) ? node.subsequent : node.else_clause
102
+ else_names = definitely_assigned_names(else_node)
103
+
104
+ (then_names & else_names).each_with_object({}) do |name, variables|
105
+ next if known_variable?(name)
106
+
107
+ declare_variable(name)
108
+ variables[name] = nil
109
+ end
110
+ end
111
+
112
+ def definitely_assigned_names(node)
113
+ return Set.new unless node
114
+
115
+ case node
116
+ when ::Prism::StatementsNode
117
+ node.body.each_with_object(Set.new) do |statement, names|
118
+ names.merge(definitely_assigned_names(statement))
119
+ end
120
+ when ::Prism::LocalVariableWriteNode, ::Prism::LocalVariableOperatorWriteNode
121
+ Set[node.name.to_sym]
122
+ when ::Prism::MultiWriteNode
123
+ Set.new(node.lefts.map { |target| target.name.to_sym })
124
+ when ::Prism::IfNode
125
+ definitely_assigned_names(node.statements) & definitely_assigned_names(node.subsequent)
126
+ when ::Prism::UnlessNode
127
+ definitely_assigned_names(node.statements) & definitely_assigned_names(node.else_clause)
128
+ when ::Prism::ElseNode
129
+ definitely_assigned_names(node.statements)
130
+ else
131
+ Set.new
132
+ end
133
+ end
134
+ end
135
+ end
136
+ end
137
+ end
@@ -0,0 +1,92 @@
1
+ # frozen_string_literal: true
2
+
3
+ module RLSL
4
+ module Prism
5
+ class ASTVisitor
6
+ module DefinitionVisiting
7
+ VISITORS = {}.tap do |visitors|
8
+ visitors[::Prism::LocalVariableWriteNode] = :visit_local_variable_write if defined?(::Prism::LocalVariableWriteNode)
9
+ visitors[::Prism::LocalVariableOperatorWriteNode] = :visit_local_variable_operator_write if defined?(::Prism::LocalVariableOperatorWriteNode)
10
+ visitors[::Prism::LocalVariableReadNode] = :visit_local_variable_read if defined?(::Prism::LocalVariableReadNode)
11
+ visitors[::Prism::DefNode] = :visit_def if defined?(::Prism::DefNode)
12
+ visitors[::Prism::GlobalVariableWriteNode] = :visit_global_variable_write if defined?(::Prism::GlobalVariableWriteNode)
13
+ visitors[::Prism::ConstantWriteNode] = :visit_constant_write if defined?(::Prism::ConstantWriteNode)
14
+ visitors[::Prism::MultiWriteNode] = :visit_multi_write if defined?(::Prism::MultiWriteNode)
15
+ visitors[::Prism::LocalVariableTargetNode] = :visit_local_variable_target if defined?(::Prism::LocalVariableTargetNode)
16
+ end.freeze
17
+
18
+ private
19
+
20
+ def visit_local_variable_write(node)
21
+ name = node.name.to_sym
22
+ value = normalize_expression(visit(node.value))
23
+ emitted_name = emitted_assignment_name(name)
24
+
25
+ if known_variable?(name)
26
+ IR::Assignment.new(IR::VarRef.new(emitted_name), value)
27
+ else
28
+ declare_variable(name)
29
+ IR::VarDecl.new(name, value)
30
+ end
31
+ end
32
+
33
+ def visit_local_variable_operator_write(node)
34
+ name = node.name.to_sym
35
+ operator = node.binary_operator.to_s
36
+ value = normalize_expression(visit(node.value))
37
+
38
+ unless known_variable?(name)
39
+ raise UnsupportedSyntaxError, "Operator assignment requires an initialized variable: #{name}"
40
+ end
41
+ emitted_name = emitted_assignment_name(name)
42
+ target = IR::VarRef.new(emitted_name)
43
+ expr = IR::BinaryOp.new(operator, IR::VarRef.new(emitted_name), IR::Parenthesized.new(value))
44
+ IR::Assignment.new(target, expr)
45
+ end
46
+
47
+ def visit_local_variable_read(node)
48
+ name = node.name.to_sym
49
+ emitted_name = fragment_parameter_reference?(name) ? emitted_parameter_name(name) : name
50
+ type = infer_param_type(name) if fragment_parameter_reference?(name)
51
+ IR::VarRef.new(emitted_name, type)
52
+ end
53
+
54
+ def visit_def(node)
55
+ params = extract_required_params(node.parameters)
56
+ body = visit_with_scoped_vars(node.body, params: params)
57
+ IR::FunctionDefinition.new(node.name.to_sym, params, body)
58
+ end
59
+
60
+ def visit_global_variable_write(node)
61
+ IR::GlobalDecl.new(node.name.to_s.sub(/^\$/, "").to_sym, visit(node.value), is_static: true)
62
+ end
63
+
64
+ def visit_constant_write(node)
65
+ IR::GlobalDecl.new(node.name.to_sym, visit(node.value), is_const: true, is_static: true)
66
+ end
67
+
68
+ def visit_multi_write(node)
69
+ if node.rest || !node.rights.empty?
70
+ raise UnsupportedSyntaxError, "Splat multiple assignment is not supported"
71
+ end
72
+
73
+ declarations = []
74
+ targets = node.lefts.map do |target|
75
+ name = target.name.to_sym
76
+ declarations << !known_variable?(name)
77
+ declare_variable(name) if declarations.last
78
+ IR::VarRef.new(emitted_assignment_name(name))
79
+ end
80
+
81
+ IR::MultipleAssignment.new(targets, visit(node.value), declarations: declarations)
82
+ end
83
+
84
+ def visit_local_variable_target(node)
85
+ name = node.name.to_sym
86
+ declare_variable(name)
87
+ IR::VarRef.new(name)
88
+ end
89
+ end
90
+ end
91
+ end
92
+ end
@@ -0,0 +1,172 @@
1
+ # frozen_string_literal: true
2
+
3
+ module RLSL
4
+ module Prism
5
+ class ASTVisitor
6
+ module ExpressionVisiting
7
+ VISITORS = {}.tap do |visitors|
8
+ visitors[::Prism::IntegerNode] = :visit_integer if defined?(::Prism::IntegerNode)
9
+ visitors[::Prism::FloatNode] = :visit_float if defined?(::Prism::FloatNode)
10
+ visitors[::Prism::RationalNode] = :visit_rational if defined?(::Prism::RationalNode)
11
+ visitors[::Prism::TrueNode] = :visit_true if defined?(::Prism::TrueNode)
12
+ visitors[::Prism::FalseNode] = :visit_false if defined?(::Prism::FalseNode)
13
+ visitors[::Prism::ParenthesesNode] = :visit_parentheses if defined?(::Prism::ParenthesesNode)
14
+ visitors[::Prism::CallNode] = :visit_call if defined?(::Prism::CallNode)
15
+ visitors[::Prism::AndNode] = :visit_and if defined?(::Prism::AndNode)
16
+ visitors[::Prism::OrNode] = :visit_or if defined?(::Prism::OrNode)
17
+ visitors[::Prism::NotNode] = :visit_not if defined?(::Prism::NotNode)
18
+ visitors[::Prism::ArrayNode] = :visit_array if defined?(::Prism::ArrayNode)
19
+ visitors[::Prism::ConstantReadNode] = :visit_constant_read if defined?(::Prism::ConstantReadNode)
20
+ visitors[::Prism::ConstantPathNode] = :visit_constant_path if defined?(::Prism::ConstantPathNode)
21
+ visitors[::Prism::GlobalVariableReadNode] = :visit_global_variable_read if defined?(::Prism::GlobalVariableReadNode)
22
+ end.freeze
23
+
24
+ private
25
+
26
+ def visit_integer(node)
27
+ IR::Literal.new(node.value, :int)
28
+ end
29
+
30
+ def visit_float(node)
31
+ IR::Literal.new(node.value, :float)
32
+ end
33
+
34
+ def visit_rational(node)
35
+ IR::Literal.new(node.value.to_f, :float)
36
+ end
37
+
38
+ def visit_true(_node)
39
+ IR::BoolLiteral.new(true)
40
+ end
41
+
42
+ def visit_false(_node)
43
+ IR::BoolLiteral.new(false)
44
+ end
45
+
46
+ def visit_parentheses(node)
47
+ inner = normalize_expression(visit(node.body))
48
+ inner = inner.statements.first if single_statement_block?(inner)
49
+ IR::Parenthesized.new(inner)
50
+ end
51
+
52
+ def visit_call(node)
53
+ return visit_call_with_block(node) if node.block
54
+
55
+ visit_plain_call(node)
56
+ end
57
+
58
+ def visit_plain_call(node)
59
+ method_name = node.name.to_s
60
+ receiver = normalize_expression(visit(node.receiver)) if node.receiver
61
+ args = node.arguments&.arguments&.map { |arg| normalize_expression(visit(arg)) } || []
62
+
63
+ if parameter_reference_call?(method_name, receiver, args)
64
+ return IR::VarRef.new(emitted_parameter_name(method_name), infer_param_type(method_name))
65
+ end
66
+ return IR::UnaryOp.new("-", receiver) if method_name == "-@" && receiver
67
+ return IR::UnaryOp.new("!", receiver) if method_name == "!" && receiver && args.empty?
68
+ return visit_receiver_call(method_name, receiver) if receiver_without_arguments?(node, receiver, args)
69
+ return IR::BinaryOp.new(method_name, receiver, args.first) if binary_operator_call?(method_name, receiver, args)
70
+ return IR::ArrayIndex.new(receiver, args.first) if method_name == "[]" && receiver && args.length == 1
71
+
72
+ IR::FuncCall.new(method_name.to_sym, args, receiver)
73
+ end
74
+
75
+ def visit_and(node)
76
+ left = normalize_expression(visit(node.left))
77
+ right = normalize_expression(visit(node.right))
78
+ IR::BinaryOp.new("&&", left, right, :bool)
79
+ end
80
+
81
+ def visit_or(node)
82
+ left = normalize_expression(visit(node.left))
83
+ right = normalize_expression(visit(node.right))
84
+ IR::BinaryOp.new("||", left, right, :bool)
85
+ end
86
+
87
+ def visit_not(node)
88
+ operand = normalize_expression(visit(node.expression))
89
+ IR::UnaryOp.new("!", operand, :bool)
90
+ end
91
+
92
+ def visit_array(node)
93
+ elements = node.elements.map { |elem| normalize_expression(visit(elem)) }
94
+ IR::ArrayLiteral.new(elements)
95
+ end
96
+
97
+ def visit_constant_read(node)
98
+ name = node.name.to_s
99
+ return IR::Constant.new(name.to_sym, :float) if %w[PI TAU].include?(name)
100
+
101
+ IR::VarRef.new(name.to_sym)
102
+ end
103
+
104
+ def visit_constant_path(node)
105
+ path_parts = []
106
+ current = node
107
+ while current.is_a?(::Prism::ConstantPathNode)
108
+ path_parts.unshift(current.name.to_s)
109
+ current = current.parent
110
+ end
111
+ path_parts.unshift(current.name.to_s) if current.respond_to?(:name)
112
+
113
+ if path_parts.first == "Math" && %w[PI TAU].include?(path_parts.last) && path_parts.length == 2
114
+ return IR::Constant.new(path_parts.last.to_sym, :float)
115
+ end
116
+
117
+ IR::VarRef.new(path_parts.join("_").to_sym)
118
+ end
119
+
120
+ def visit_global_variable_read(node)
121
+ IR::VarRef.new(node.name.to_s.sub(/^\$/, "").to_sym)
122
+ end
123
+
124
+ def normalize_expression(node)
125
+ return node unless node.is_a?(IR::IfStatement)
126
+
127
+ then_expr = extract_if_branch_expr(node.then_branch)
128
+ else_expr = extract_if_branch_expr(node.else_branch)
129
+ unless then_expr && else_expr
130
+ raise UnsupportedSyntaxError, "Unsupported if-expression: branches must be single expressions"
131
+ end
132
+
133
+ IR::Ternary.new(node.condition, then_expr, else_expr).tap do |ternary|
134
+ ternary.location = node.location
135
+ end
136
+ end
137
+
138
+ def extract_if_branch_expr(branch)
139
+ return nil unless branch
140
+ return branch unless branch.is_a?(IR::Block)
141
+ return nil if branch.statements.empty? || branch.statements.length != 1
142
+
143
+ branch.statements.first
144
+ end
145
+
146
+ def single_statement_block?(node)
147
+ node.is_a?(IR::Block) && node.statements.length == 1
148
+ end
149
+
150
+ def parameter_reference_call?(method_name, receiver, args)
151
+ !receiver && args.empty? &&
152
+ (parameter_reference?(method_name.to_sym) || infer_param_type(method_name.to_sym))
153
+ end
154
+
155
+ def receiver_without_arguments?(node, receiver, args)
156
+ receiver && args.empty? && !node.arguments
157
+ end
158
+
159
+ def visit_receiver_call(method_name, receiver)
160
+ return IR::FieldAccess.new(receiver, method_name) if receiver.type == :uniforms
161
+ return receiver if method_name == "freeze"
162
+
163
+ IR::FieldAccess.new(receiver, method_name)
164
+ end
165
+
166
+ def binary_operator_call?(method_name, receiver, args)
167
+ BINARY_OPERATORS.include?(method_name) && receiver && args.length == 1
168
+ end
169
+ end
170
+ end
171
+ end
172
+ end
@@ -0,0 +1,48 @@
1
+ # frozen_string_literal: true
2
+
3
+ module RLSL
4
+ module Prism
5
+ class ASTVisitor
6
+ class ScopeContext
7
+ Scope = Struct.new(:params, :declared_vars)
8
+
9
+ def initialize(params: [])
10
+ @scopes = [build_scope(params)]
11
+ end
12
+
13
+ def with_scope(params: [])
14
+ @scopes << build_scope(params)
15
+ yield
16
+ ensure
17
+ @scopes.pop
18
+ end
19
+
20
+ def parameter?(name)
21
+ @scopes.reverse_each.any? { |scope| scope.params.include?(name.to_sym) }
22
+ end
23
+
24
+ def root_parameter?(name)
25
+ @scopes.rindex { |scope| scope.params.include?(name.to_sym) } == 0
26
+ end
27
+
28
+ def declared?(name)
29
+ @scopes.reverse_each.any? { |scope| scope.declared_vars.include?(name.to_sym) }
30
+ end
31
+
32
+ def known_variable?(name)
33
+ parameter?(name) || declared?(name)
34
+ end
35
+
36
+ def declare(name)
37
+ @scopes.last.declared_vars.add(name.to_sym)
38
+ end
39
+
40
+ private
41
+
42
+ def build_scope(params)
43
+ Scope.new(Set.new(Array(params).map(&:to_sym)), Set.new)
44
+ end
45
+ end
46
+ end
47
+ end
48
+ end
@@ -0,0 +1,21 @@
1
+ # frozen_string_literal: true
2
+
3
+ module RLSL
4
+ module Prism
5
+ class ASTVisitor
6
+ class VisitorRegistry
7
+ TRANSPARENT_NODES = [
8
+ (::Prism::ArgumentsNode if defined?(::Prism::ArgumentsNode)),
9
+ (::Prism::BlockParametersNode if defined?(::Prism::BlockParametersNode)),
10
+ (::Prism::ParametersNode if defined?(::Prism::ParametersNode))
11
+ ].compact.freeze
12
+
13
+ def self.build(*visitor_maps)
14
+ visitor_maps.each_with_object({}) do |visitor_map, merged|
15
+ merged.merge!(visitor_map)
16
+ end.freeze
17
+ end
18
+ end
19
+ end
20
+ end
21
+ end