46日目: 1行 pattern matchingに対応する
前回はAs patternに対応をしました。 今回は1行 pattern matchingに取り組みます。
1行 pattern matching
1行 pattern matchingには2つの異なる書き方があります。
inによるpattern matchingの場合、patternにマッチすればtrueを、マッチしなければfalseを返します。
一方で=>によるpattern matchingの場合、マッチしないときはNoMatchingPatternError例外が投げられます。
a in [0, 1, 2] b => [0, 1, 2] #=> []: [] length mismatch (given 0, expected 3) (NoMatchingPatternError)
対応するノードについてみておきましょう。
まずはa in [0, 1, 2]からです。
# @ NODE_CASE3 (id: 11, line: 1, location: (1,0)-(1,14))* # +- nd_head: # | @ NODE_VCALL (id: 0, line: 1, location: (1,0)-(1,1)) # | +- nd_mid: :a # +- nd_body: # | @ NODE_IN (id: 10, line: 1, location: (1,5)-(1,14)) # | +- nd_head: # | | @ NODE_ARYPTN (id: 7, line: 1, location: (1,6)-(1,13)) # | | +- nd_pconst: # | | | (null node) # | | +- pre_args: # | | | @ NODE_LIST (id: 2, line: 1, location: (1,6)-(1,13)) # | | | +- as.nd_alen: 3 # | | | +- nd_head: # | | | | @ NODE_INTEGER (id: 1, line: 1, location: (1,6)-(1,7)) # | | | | +- val: 0 # | | | +- nd_head: # | | | | @ NODE_INTEGER (id: 3, line: 1, location: (1,9)-(1,10)) # | | | | +- val: 1 # | | | +- nd_head: # | | | | @ NODE_INTEGER (id: 5, line: 1, location: (1,12)-(1,13)) # | | | | +- val: 2 # | | | +- nd_next: # | | | (null node) # | | +- rest_arg: # | | | (null node) # | | +- post_args: # | | (null node) # | +- nd_body: # | | @ NODE_TRUE (id: 8, line: 1, location: (1,5)-(1,14)) # | +- nd_next: # | | @ NODE_FALSE (id: 9, line: 1, location: (1,5)-(1,14))
変更前はNODE_CASE3とNODE_INで表現しています。
そしてNODE_INのnd_bodyとnd_nextがそれぞれNODE_TRUEとNODE_FALSEになっています。
これは以下のコードと同じノードになっています。
case a in [0, 1, 2] true else false end
書き換え後はというとMatchPredicateNodeという専用のノードを使って表現するようになります。
# @ MatchPredicateNode (location: (1,0)-(1,14)) # +-- value: # | @ CallNode (location: (1,0)-(1,1)) # | +-- name: :a # +-- pattern: # | @ ArrayPatternNode (location: (1,5)-(1,14)) # | +-- constant: nil # | +-- requireds: (length: 3) # | | +-- @ IntegerNode (location: (1,6)-(1,7)) # | | | +-- IntegerBaseFlags: decimal # | | | +-- value: 0 # | | +-- @ IntegerNode (location: (1,9)-(1,10)) # | | | +-- IntegerBaseFlags: decimal # | | | +-- value: 1 # | | +-- @ IntegerNode (location: (1,12)-(1,13)) # | | +-- IntegerBaseFlags: decimal # | | +-- value: 2 # | +-- rest: nil # | +-- posts: (length: 0)
つぎにa => [0, 1, 2]の場合です。
# @ NODE_CASE3 (id: 9, line: 1, location: (1,0)-(1,14))* # +- nd_head: # | @ NODE_VCALL (id: 0, line: 1, location: (1,0)-(1,1)) # | +- nd_mid: :b # +- nd_body: # | @ NODE_IN (id: 8, line: 1, location: (1,5)-(1,14)) # | +- nd_head: # | | @ NODE_ARYPTN (id: 7, line: 1, location: (1,6)-(1,13)) # | | +- nd_pconst: # | | | (null node) # | | +- pre_args: # | | | @ NODE_LIST (id: 2, line: 1, location: (1,6)-(1,13)) # | | | +- as.nd_alen: 3 # | | | +- nd_head: # | | | | @ NODE_INTEGER (id: 1, line: 1, location: (1,6)-(1,7)) # | | | | +- val: 0 # | | | +- nd_head: # | | | | @ NODE_INTEGER (id: 3, line: 1, location: (1,9)-(1,10)) # | | | | +- val: 1 # | | | +- nd_head: # | | | | @ NODE_INTEGER (id: 5, line: 1, location: (1,12)-(1,13)) # | | | | +- val: 2 # | | | +- nd_next: # | | | (null node) # | | +- rest_arg: # | | | (null node) # | | +- post_args: # | | (null node) # | +- nd_body: # | | (null node) # | +- nd_next: # | | (null node)
書き換え前はinの場合と同様にNODE_CASE3とNODE_INで表現しています。
inと異なりnd_bodyとnd_nextがnullになっています。
これは以下のコードとおおよそ同じノードになっています。
case a in [0, 1, 2] end
書き換え後はというとMatchRequiredNodeという専用のノードを使って表現するようになります。
# @ MatchRequiredNode (location: (1,0)-(1,14)) # +-- value: # | @ CallNode (location: (1,0)-(1,1)) # | +-- name: :b # +-- pattern: # | @ ArrayPatternNode (location: (1,5)-(1,14)) # | +-- constant: nil # | +-- requireds: (length: 3) # | | +-- @ IntegerNode (location: (1,6)-(1,7)) # | | | +-- IntegerBaseFlags: decimal # | | | +-- value: 0 # | | +-- @ IntegerNode (location: (1,9)-(1,10)) # | | | +-- IntegerBaseFlags: decimal # | | | +-- value: 1 # | | +-- @ IntegerNode (location: (1,12)-(1,13)) # | | +-- IntegerBaseFlags: decimal # | | +-- value: 2 # | +-- rest: nil # | +-- posts: (length: 0)
parse.yを書き換える
parse.yの書き換えはシンプルで、対応する生成規則のアクションが生成するノードの種類を変更するだけです。
@@ -4001,7 +4005,7 @@ expr : command_call p->ctxt.in_kwarg = $ctxt.in_kwarg; p->ctxt.in_alt_pattern = $ctxt.in_alt_pattern; p->ctxt.capture_in_pattern = $ctxt.capture_in_pattern; - $$ = NEW_CASE3($arg, NEW_IN($body, 0, 0, &@body, &NULL_LOC, &NULL_LOC, &@2), &@$, &NULL_LOC, &NULL_LOC); + $$ = NEW_RB_MATCH_REQUIRED($arg, $body, &@$, &@tASSOC); /*% ripper: case!($:arg, in!($:body, Qnil, Qnil)) %*/ } | arg keyword_in @@ -4016,7 +4020,7 @@ expr : command_call p->ctxt.in_kwarg = $ctxt.in_kwarg; p->ctxt.in_alt_pattern = $ctxt.in_alt_pattern; p->ctxt.capture_in_pattern = $ctxt.capture_in_pattern; - $$ = NEW_CASE3($arg, NEW_IN($body, NEW_RB_TRUE(&@body), NEW_RB_FALSE(&@body), &@body, &@keyword_in, &NULL_LOC, &NULL_LOC), &@$, &NULL_LOC, &NULL_LOC); + $$ = NEW_RB_MATCH_PREDICATE($arg, $body, &@$, &@keyword_in); /*% ripper: case!($:arg, in!($:body, Qnil, Qnil)) %*/ }
compile.cを書き換える
いままではcase v in ...もa in [0, 1, 2]もb => [0, 1, 2]も同じ種類のノードで表現していました。
そのためcompile.cのcompile_case3という関数に任せることができました。
今回ノードを分割したため、case v in ...とその他ではpatternが単数か複数かという差異が生じるようになりました。
a in [0, 1, 2]やb => [0, 1, 2]はcase v in ...の特殊形なのですが、ここでは下手に関数を共通化せずにMatchPredicateNodeとMatchRequiredNode用にそれぞれ専用の関数を用意することにします。
それぞれのケースのバイトコードを確認するところから始めましょう。
まずはMatchPredicateNodeのケースです。
# 0000 putnil ( 1)[Li] # 0001 putself # 0002 opt_send_without_block <calldata!mid:a, argc:0, FCALL|VCALL|ARGS_SIMPLE> # 0004 dup # 0005 topn 2 # 0007 branchnil 18 # 0009 topn 2 # 0011 branchunless 86 # 0013 pop # 0014 topn 1 # 0016 jump 36 # 0018 dup # 0019 putobject :deconstruct # 0021 opt_send_without_block <calldata!mid:respond_to?, argc:1, ARGS_SIMPLE> # 0023 setn 3 # 0025 branchunless 86 # 0027 opt_send_without_block <calldata!mid:deconstruct, argc:0, ARGS_SIMPLE> # 0029 setn 2 # 0031 dup # 0032 checktype T_ARRAY # 0034 branchunless 77 # 0036 dup # 0037 opt_length <calldata!mid:length, argc:0, ARGS_SIMPLE>[CcCr] # 0039 putobject 3 # 0041 opt_eq <calldata!mid:==, argc:1, ARGS_SIMPLE>[CcCr] # 0043 branchunless 86 # 0045 dup # 0046 putobject_INT2FIX_0_ # 0047 opt_aref <calldata!mid:[], argc:1, ARGS_SIMPLE>[CcCr] # 0049 putobject_INT2FIX_0_ # 0050 checkmatch 2 # 0052 branchunless 86 # 0054 dup # 0055 putobject_INT2FIX_1_ # 0056 opt_aref <calldata!mid:[], argc:1, ARGS_SIMPLE>[CcCr] # 0058 putobject_INT2FIX_1_ # 0059 checkmatch 2 # 0061 branchunless 86 # 0063 dup # 0064 putobject 2 # 0066 opt_aref <calldata!mid:[], argc:1, ARGS_SIMPLE>[CcCr] # 0068 putobject 2 # 0070 checkmatch 2 # 0072 branchunless 86 # 0074 pop # 0075 jump 92 # 0077 putspecialobject 1 # 0079 putobject TypeError # 0081 putobject "deconstruct must return Array" # 0083 opt_send_without_block <calldata!mid:core#raise, argc:2, ARGS_SIMPLE> # 0085 pop # 0086 pop # 0087 pop # 0088 pop # 0089 putobject false # 0091 leave # 0092 adjuststack 2 # 0094 putobject true # 0096 leave a in [0, 1, 2]
さきほど触れたように以下のコードと同様のバイトコードになります。
そのためマッチに成功したときの値として0094 putobject trueが、失敗したときの値として0089 putobject falseのバイトコードが生成されています。
case a in [0, 1, 2] true else false end
compile_match_predicate関数の実装はcompile_case3関数の実装から必要な部分を切り出したものになります。
static int compile_match_predicate(rb_iseq_t *iseq, LINK_ANCHOR *const ret, const rb_match_predicate_node_t *const node, int popped) { const NODE *line_node = node->pattern; LABEL *endlabel, *elselabel; DECL_ANCHOR(head); DECL_ANCHOR(body_seq); DECL_ANCHOR(cond_seq); int line; VALUE branches = 0; int branch_id = 0; INIT_ANCHOR(head); INIT_ANCHOR(body_seq); INIT_ANCHOR(cond_seq); branches = decl_branch_base(iseq, PTR2NUM(node), nd_code_loc(node), "case"); line = nd_line(line_node); endlabel = NEW_LABEL(line); elselabel = NEW_LABEL(line); ADD_INSN(head, line_node, putnil); /* allocate stack for cached #deconstruct value */ CHECK(COMPILE(head, "case base", node->value)); ADD_SEQ(ret, head); /* case VAL */ { const NODE *pattern = node->pattern; line = nd_line(pattern); LABEL *l1 = NEW_LABEL(line); ADD_LABEL(body_seq, l1); ADD_INSN1(body_seq, line_node, adjuststack, INT2FIX(2)); const NODE *const coverage_node = pattern; add_trace_branch_coverage( iseq, body_seq, nd_code_loc(coverage_node), nd_node_id(coverage_node), branch_id++, "in", branches); if (!popped) ADD_INSN1(body_seq, line_node, putobject, Qtrue); ADD_INSNL(body_seq, line_node, jump, endlabel); int pat_line = nd_line(pattern); LABEL *next_pat = NEW_LABEL(pat_line); ADD_INSN (cond_seq, pattern, dup); /* dup case VAL */ // NOTE: set base_index (it's "under" the matchee value, so it's position is 2) CHECK(iseq_compile_pattern_each(iseq, cond_seq, pattern, l1, next_pat, false, false, 2, true)); ADD_LABEL(cond_seq, next_pat); LABEL_UNREMOVABLE(next_pat); } { ADD_LABEL(cond_seq, elselabel); ADD_INSN(cond_seq, line_node, pop); ADD_INSN(cond_seq, line_node, pop); /* discard cached #deconstruct value */ add_trace_branch_coverage(iseq, cond_seq, nd_code_loc(node), nd_node_id(node), branch_id, "else", branches); if (!popped) ADD_INSN1(cond_seq, line_node, putobject, Qfalse); ADD_INSNL(cond_seq, line_node, jump, endlabel); ADD_INSN(cond_seq, line_node, putnil); if (popped) { ADD_INSN(cond_seq, line_node, putnil); } } ADD_SEQ(ret, cond_seq); ADD_SEQ(ret, body_seq); ADD_LABEL(ret, endlabel); return COMPILE_OK; }
buildして実行してみると期待した通りの結果が得られます。
def m(a) a in [0, 1, 2] end p m([]) #=> false p m([0, 1, 2]) #=> true
次にMatchRequiredNodeのケースです。
# 0000 putnil ( 1)[Li] # 0001 putnil # 0002 putobject false # 0004 putnil # 0005 putnil # 0006 putself # 0007 opt_send_without_block <calldata!mid:b, argc:0, FCALL|VCALL|ARGS_SIMPLE> # 0009 dup # 0010 topn 2 # 0012 branchnil 23 # 0014 topn 2 # 0016 branchunless 215 # 0018 pop # 0019 topn 1 # 0021 jump 60 # 0023 dup # 0024 putobject :deconstruct # 0026 opt_send_without_block <calldata!mid:respond_to?, argc:1, ARGS_SIMPLE> # 0028 setn 3 # 0030 dup # 0031 branchif 49 # 0033 putspecialobject 1 # 0035 putobject "%p does not respond to #deconstruct" # 0037 topn 3 # 0039 opt_send_without_block <calldata!mid:core#sprintf, argc:2, ARGS_SIMPLE> # 0041 setn 5 # 0043 putobject false # 0045 setn 7 # 0047 pop # 0048 pop # 0049 branchunless 215 # 0051 opt_send_without_block <calldata!mid:deconstruct, argc:0, ARGS_SIMPLE> # 0053 setn 2 # 0055 dup # 0056 checktype T_ARRAY # 0058 branchunless 206 # 0060 dup # 0061 opt_length <calldata!mid:length, argc:0, ARGS_SIMPLE>[CcCr] # 0063 putobject 3 # 0065 opt_eq <calldata!mid:==, argc:1, ARGS_SIMPLE>[CcCr] # 0067 dup # 0068 branchif 91 # 0070 putspecialobject 1 # 0072 putobject "%p length mismatch (given %p, expected %p)" # 0074 topn 3 # 0076 dup # 0077 opt_length <calldata!mid:length, argc:0, ARGS_SIMPLE>[CcCr] # 0079 putobject 3 # 0081 opt_send_without_block <calldata!mid:core#sprintf, argc:4, ARGS_SIMPLE> # 0083 setn 5 # 0085 putobject false # 0087 setn 7 # 0089 pop # 0090 pop # 0091 branchunless 215 # 0093 dup # 0094 putobject_INT2FIX_0_ # 0095 opt_aref <calldata!mid:[], argc:1, ARGS_SIMPLE>[CcCr] # 0097 putobject_INT2FIX_0_ # 0098 dupn 2 # 0100 checkmatch 2 # 0102 dup # 0103 branchif 123 # 0105 putspecialobject 1 # 0107 putobject "%p === %p does not return true" # 0109 topn 3 # 0111 topn 5 # 0113 opt_send_without_block <calldata!mid:core#sprintf, argc:3, ARGS_SIMPLE> # 0115 setn 7 # 0117 putobject false # 0119 setn 9 # 0121 pop # 0122 pop # 0123 setn 2 # 0125 pop # 0126 pop # 0127 branchunless 215 # 0129 dup # 0130 putobject_INT2FIX_1_ # 0131 opt_aref <calldata!mid:[], argc:1, ARGS_SIMPLE>[CcCr] # 0133 putobject_INT2FIX_1_ # 0134 dupn 2 # 0136 checkmatch 2 # 0138 dup # 0139 branchif 159 # 0141 putspecialobject 1 # 0143 putobject "%p === %p does not return true" # 0145 topn 3 # 0147 topn 5 # 0149 opt_send_without_block <calldata!mid:core#sprintf, argc:3, ARGS_SIMPLE> # 0151 setn 7 # 0153 putobject false # 0155 setn 9 # 0157 pop # 0158 pop # 0159 setn 2 # 0161 pop # 0162 pop # 0163 branchunless 215 # 0165 dup # 0166 putobject 2 # 0168 opt_aref <calldata!mid:[], argc:1, ARGS_SIMPLE>[CcCr] # 0170 putobject 2 # 0172 dupn 2 # 0174 checkmatch 2 # 0176 dup # 0177 branchif 197 # 0179 putspecialobject 1 # 0181 putobject "%p === %p does not return true" # 0183 topn 3 # 0185 topn 5 # 0187 opt_send_without_block <calldata!mid:core#sprintf, argc:3, ARGS_SIMPLE> # 0189 setn 7 # 0191 putobject false # 0193 setn 9 # 0195 pop # 0196 pop # 0197 setn 2 # 0199 pop # 0200 pop # 0201 branchunless 215 # 0203 pop # 0204 jump 262 # 0206 putspecialobject 1 # 0208 putobject TypeError # 0210 putobject "deconstruct must return Array" # 0212 opt_send_without_block <calldata!mid:core#raise, argc:2, ARGS_SIMPLE> # 0214 pop # 0215 pop # 0216 putspecialobject 1 # 0218 topn 4 # 0220 branchif 238 # 0222 putobject NoMatchingPatternError # 0224 putspecialobject 1 # 0226 putobject "%p: %s" # 0228 topn 4 # 0230 topn 7 # 0232 opt_send_without_block <calldata!mid:core#sprintf, argc:3, ARGS_SIMPLE> # 0234 opt_send_without_block <calldata!mid:core#raise, argc:2, ARGS_SIMPLE> # 0236 jump 258 # 0238 putobject NoMatchingPatternKeyError # 0240 putspecialobject 1 # 0242 putobject "%p: %s" # 0244 topn 4 # 0246 topn 7 # 0248 opt_send_without_block <calldata!mid:core#sprintf, argc:3, ARGS_SIMPLE> # 0250 topn 7 # 0252 topn 9 # 0254 opt_send_without_block <calldata!mid:new, argc:3, kw:[#<Symbol:0x000000000023410c>,#<Symbol:0x000000000022010c>], KWARG> # 0256 opt_send_without_block <calldata!mid:core#raise, argc:1, ARGS_SIMPLE> # 0258 adjuststack 7 # 0260 putnil # 0261 leave # 0262 adjuststack 6 # 0264 putnil # 0265 leave b => [0, 1, 2]
0216 putspecialobject 1から0256 opt_send_without_block <calldata!mid:core#raise, argc:1, ARGS_SIMPLE>がmatchしなかったときにNoMatchingPatternErrorを投げるためのバイトコードです。
matchに成功した場合にはnilになるというのは、以下のバイトコードが対応しています。
# 0262 adjuststack 6 # 0264 putnil # 0265 leave
compile_match_required関数の実装もまたcompile_case3関数の実装から必要な部分を抜き出したような実装になりますが、ここでは割愛します。
buildして実行してみると期待した通りの結果が得られます。
def m(a) a => [0, 1, 2] end p m([0, 1, 2]) #=> nil p m([]) #=> []: [] length mismatch (given 0, expected 3) (NoMatchingPatternError)
まとめ
- 1行 pattern matchingに対応した
パターンマッチング全体の進捗は以下の通りです。
Value patternp_primitive ("str",1,:symなど)range_expr (1...3など)p_var_ref (^varなど)p_expr_ref (^(cmd 1, 2)など)p_const (A,::A,A::Bなど)
Variable patternArray patternHash patternFind patternAlternative patternAs pattern後置ifと後置unless1行 pattern matching
パターンマッチングが終わったので次回はforに取り組む予定です。
