CoolFace
Apppublic

idsedykh/codebleu

sourceHugging Faceupdated 4y agoView on Hugging Face
0likes
dataflow_match.py144 linesDownload Raw Back to root
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