idsedykh/codebleu
0
1# Copyright (c) Microsoft Corporation.
2# Licensed under the MIT license.
3
4import os
5from .parser import DFG_python,DFG_java,DFG_ruby,DFG_go,DFG_php,DFG_javascript,DFG_csharp
6from .parser import (remove_comments_and_docstrings,
7 tree_to_token_index,
8 index_to_code_token,
9 tree_to_variable_index)
10from tree_sitter import Language, Parser
11import pdb
12
13dfg_function={
14 'python':DFG_python,
15 'java':DFG_java,
16 'ruby':DFG_ruby,
17 'go':DFG_go,
18 'php':DFG_php,
19 'javascript':DFG_javascript,
20 'c_sharp':DFG_csharp,
21}
22
23def calc_dataflow_match(references, candidate, lang):
24 return corpus_dataflow_match([references], [candidate], lang)
25
26def corpus_dataflow_match(references, candidates, lang):
27 LANGUAGE = Language(os.path.abspath(os.path.dirname(__file__)) + '/parser/my-languages.so', lang)
28 parser = Parser()
29 parser.set_language(LANGUAGE)
30 parser = [parser,dfg_function[lang]]
31 match_count = 0
32 total_count = 0
33
34 for i in range(len(candidates)):
35 references_sample = references[i]
36 candidate = candidates[i]
37 for reference in references_sample:
38 try:
39 candidate=remove_comments_and_docstrings(candidate,'java')
40 except:
41 pass
42 try:
43 reference=remove_comments_and_docstrings(reference,'java')
44 except:
45 pass
46
47 cand_dfg = get_data_flow(candidate, parser)
48 ref_dfg = get_data_flow(reference, parser)
49
50 normalized_cand_dfg = normalize_dataflow(cand_dfg)
51 normalized_ref_dfg = normalize_dataflow(ref_dfg)
52
53 if len(normalized_ref_dfg) > 0:
54 total_count += len(normalized_ref_dfg)
55 for dataflow in normalized_ref_dfg:
56 if dataflow in normalized_cand_dfg:
57 match_count += 1
58 normalized_cand_dfg.remove(dataflow)
59 if total_count == 0:
60 print("WARNING: There is no reference data-flows extracted from the whole corpus, and the data-flow match score degenerates to 0. Please consider ignoring this score.")
61 return 0
62 score = match_count / total_count
63 return score
64
65def get_data_flow(code, parser):
66 try:
67 tree = parser[0].parse(bytes(code,'utf8'))
68 root_node = tree.root_node
69 tokens_index=tree_to_token_index(root_node)
70 code=code.split('\n')
71 code_tokens=[index_to_code_token(x,code) for x in tokens_index]
72 index_to_code={}
73 for idx,(index,code) in enumerate(zip(tokens_index,code_tokens)):
74 index_to_code[index]=(idx,code)
75 try:
76 DFG,_=parser[1](root_node,index_to_code,{})
77 except:
78 DFG=[]
79 DFG=sorted(DFG,key=lambda x:x[1])
80 indexs=set()
81 for d in DFG:
82 if len(d[-1])!=0:
83 indexs.add(d[1])
84 for x in d[-1]:
85 indexs.add(x)
86 new_DFG=[]
87 for d in DFG:
88 if d[1] in indexs:
89 new_DFG.append(d)
90 codes=code_tokens
91 dfg=new_DFG
92 except:
93 codes=code.split()
94 dfg=[]
95 #merge nodes
96 dic={}
97 for d in dfg:
98 if d[1] not in dic:
99 dic[d[1]]=d
100 else:
101 dic[d[1]]=(d[0],d[1],d[2],list(set(dic[d[1]][3]+d[3])),list(set(dic[d[1]][4]+d[4])))
102 DFG=[]
103 for d in dic:
104 DFG.append(dic[d])
105 dfg=DFG
106 return dfg
107
108def normalize_dataflow_item(dataflow_item):
109 var_name = dataflow_item[0]
110 var_pos = dataflow_item[1]
111 relationship = dataflow_item[2]
112 par_vars_name_list = dataflow_item[3]
113 par_vars_pos_list = dataflow_item[4]
114
115 var_names = list(set(par_vars_name_list+[var_name]))
116 norm_names = {}
117 for i in range(len(var_names)):
118 norm_names[var_names[i]] = 'var_'+str(i)
119
120 norm_var_name = norm_names[var_name]
121 relationship = dataflow_item[2]
122 norm_par_vars_name_list = [norm_names[x] for x in par_vars_name_list]
123
124 return (norm_var_name, relationship, norm_par_vars_name_list)
125
126def normalize_dataflow(dataflow):
127 var_dict = {}
128 i = 0
129 normalized_dataflow = []
130 for item in dataflow:
131 var_name = item[0]
132 relationship = item[2]
133 par_vars_name_list = item[3]
134 for name in par_vars_name_list:
135 if name not in var_dict:
136 var_dict[name] = 'var_'+str(i)
137 i += 1
138 if var_name not in var_dict:
139 var_dict[var_name] = 'var_'+str(i)
140 i+= 1
141 normalized_dataflow.append((var_dict[var_name], relationship, [var_dict[x] for x in par_vars_name_list]))
142 return normalized_dataflow
143
144 