JetBrains-Research/commit-message-editing
0
1import json2import os3import random4import uuid5from datetime import datetime6from difflib import ndiff7 8import gradio as gr9 10from data_loader import load_data11from hf_dataset_saver_builder import get_dataset_saver12 13HF_TOKEN = os.environ.get('HF_REWRITING_TOKEN')14HF_DATASET = os.environ.get('HF_REWRITING_DATASET')15 16data = load_data()17 18n_samples = len(data)19 20saver = get_dataset_saver(HF_TOKEN, HF_DATASET, private=True, separate_dirs=True)21 22 23def convert_diff_to_unified(diff_string):24 diff = json.loads(diff_string)25 26 result = "\n".join(27 [28 f'--- {modified_file["old_path"]}\n'29 f'+++ {modified_file["new_path"]}\n'30 f'{modified_file["diff"]}'31 for modified_file in diff32 ]33 )34 35 return result36 37 38def get_diff2html_view(raw_diff):39 html = f"""40 <div style='width:100%; height:1400px; overflow:auto; position: relative'>41 <div id='diff-raw' hidden>{raw_diff}</div> 42 <div class="d2h-view-wrapper">43 <div id='diff-view'></div>44 </div>45 </div>46 """47 48 return html49 50 51def get_github_link_md(repo, hash):52 return f'[See the commit on Github](https://github.com/{repo}/commit/{hash})'53 54 55def char_diff_obj(change_type, pos, character, timestamp):56 return {"t": change_type, "p": pos, "c": character, "ts": timestamp}57 58 59def update_commit_view(sample_ind):60 if sample_ind >= n_samples:61 return None62 63 record = data[sample_ind]64 65 diff_view = get_diff2html_view(convert_diff_to_unified(record['mods']))66 67 repo_val = record['repo']68 hash_val = record['hash']69 github_link_md = get_github_link_md(repo_val, hash_val)70 71 diff_loaded_timestamp = datetime.now().isoformat()72 73 summary_md = f"{record['summary']}"74 75 commit_message = record['prediction']76 commit_message_start = commit_message77 commit_message_prev = commit_message78 commit_message_history = []79 80 return (81 github_link_md, diff_view, repo_val, hash_val, diff_loaded_timestamp, summary_md,82 commit_message_start, commit_message, commit_message_prev, commit_message_history)83 84 85def next_sample(current_sample_ind, shuffled_idx):86 if current_sample_ind == n_samples:87 return None88 89 current_sample_ind += 190 updated_view = update_commit_view(shuffled_idx[current_sample_ind])91 return (current_sample_ind,) + updated_view92 93 94with open("head.html") as head_file:95 head_html = head_file.read()96 97force_light_theme_js_func = """98function refresh() {99 const url = new URL(window.location);100 101 if (url.searchParams.get('__theme') !== 'light') {102 url.searchParams.set('__theme', 'light');103 window.location.href = url.href;104 }105}106"""107 108with gr.Blocks(theme=gr.themes.Soft(), head=head_html, css="style_overrides.css",109 js=force_light_theme_js_func) as application:110 repo_val = gr.Textbox(interactive=False, label='repo', visible=False)111 hash_val = gr.Textbox(interactive=False, label='hash', visible=False)112 shuffled_idx_val = gr.JSON(visible=False)113 114 with gr.Row():115 with gr.Accordion("Help"):116 with open("survey_guide.md") as content_file:117 gr.Markdown(content_file.read())118 119 with gr.Row():120 current_sample_sld = gr.Slider(minimum=0, maximum=n_samples, step=1,121 value=0,122 interactive=False,123 label='sample_ind',124 info=f"Samples labeled/skipped",125 show_label=False,126 container=False,127 scale=5)128 129 with gr.Column(scale=1):130 gr.Markdown(value=f"#### Total number of samples: {n_samples}")131 with gr.Column(scale=1):132 skip_btn = gr.Button("Skip the current sample")133 with gr.Row():134 with gr.Column(scale=2):135 github_link = gr.Markdown()136 diff_view = gr.HTML()137 with gr.Column(scale=1):138 with gr.Accordion("Commit summary (AI generated)", open=False):139 commit_summary = gr.Markdown()140 commit_msg_start = gr.TextArea(label="commit_msg_start", visible=False)141 142 gr.Markdown(value=f"#### Please, edit the message in the text box below")143 commit_msg = gr.TextArea(label="commit_msg_end", show_label=False,144 info="Commit message (can be scrollable)")145 commit_msg_prev = gr.TextArea(visible=False)146 commit_msg_history = gr.JSON(label="commit_msg_history", visible=False)147 148 submit_btn = gr.Button("Submit")149 150 session_val = gr.Textbox(info='Session', interactive=False, container=True, show_label=False,151 label='session')152 153 with gr.Row(visible=False):154 sample_loaded_timestamp = gr.Textbox(info="Sample loaded", label='loaded_ts', interactive=False,155 container=True, show_label=False)156 now_timestamp = gr.Textbox(info="Current time",157 interactive=False, container=True, show_label=False,158 value=lambda: datetime.now().isoformat(), every=0.1,159 label='submitted_ts')160 161 commit_view = [162 github_link,163 diff_view,164 repo_val,165 hash_val,166 sample_loaded_timestamp,167 commit_summary,168 commit_msg_start,169 commit_msg,170 commit_msg_prev,171 commit_msg_history172 ]173 174 feedback_metadata = [175 session_val,176 repo_val,177 hash_val,178 sample_loaded_timestamp,179 now_timestamp180 ]181 182 feedback_form = [183 commit_msg_start,184 commit_msg,185 commit_msg_history186 ]187 188 saver.setup([current_sample_sld] + feedback_metadata + feedback_form, "feedback")189 190 skip_btn.click(next_sample, inputs=[current_sample_sld, shuffled_idx_val],191 outputs=[current_sample_sld] + commit_view)192 193 194 def submit(current_sample, shuffled_idx, *args):195 saver.flag((current_sample,) + args)196 return next_sample(current_sample, shuffled_idx)197 198 199 submit_btn.click(200 submit,201 inputs=[current_sample_sld, shuffled_idx_val] + feedback_metadata + feedback_form,202 outputs=[current_sample_sld] + commit_view203 )204 205 206 def on_commit_msg_changed(message, prev_message, history):207 timestamp = datetime.now().isoformat()208 for i, s in enumerate(ndiff(prev_message, message)):209 diff = char_diff_obj(s[0], i, s[-1], timestamp)210 if diff['t'] in ('+', '-'):211 history.append(diff)212 return message, history213 214 215 commit_msg.change(on_commit_msg_changed, inputs=[commit_msg, commit_msg_prev, commit_msg_history],216 outputs=[commit_msg_prev, commit_msg_history])217 218 219 def init_session(current_sample):220 session = str(uuid.uuid4())221 shuffled_idx = list(range(n_samples))222 random.shuffle(shuffled_idx)223 return (session, shuffled_idx) + update_commit_view(shuffled_idx[current_sample])224 225 226 application.load(init_session,227 inputs=[current_sample_sld],228 outputs=[session_val, shuffled_idx_val] + commit_view, )229 230application.launch()231 