Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
54 changes: 47 additions & 7 deletions backend/agents.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,12 @@
import re

import time
from typing import Callable

from openai import OpenAI

from modules.exceptions import GenerationCancelled


class Agent:
def __init__(self, llm_name: str):
Expand All @@ -19,9 +22,17 @@ def __init__(self, llm_name: str):
self.top_p = 0.1
self.seed = 1203
self.max_completion_tokens = 5120
self.cancel_check: Callable[[], bool] = lambda: False

def get_response(self, messages, n=1, skip_deepseek_think: bool=False):
if self.model_name == 'gpt-4o' or self.model_name == 'gpt-3.5-turbo':
self._check_cancel()
if self.model_name in (
'gpt-4o',
'gpt-3.5-turbo',
'qwen-plus',
'qwen-coder-plus',
'qwen-long-latest',
):
if self.system_prompt:
messages = [{'role': 'system', 'content': self.system_prompt}] + messages
response = self._get_gpt_response(messages, n=n)
Expand All @@ -41,12 +52,23 @@ def get_response(self, messages, n=1, skip_deepseek_think: bool=False):
else:
raise ValueError(f"Unknown LLM name: {self.model_name}")
return response

def set_cancel_check(self, checker: Callable[[], bool] | None) -> None:
if checker:
self.cancel_check = checker
else:
self.cancel_check = lambda: False

def _check_cancel(self) -> None:
if self.cancel_check and self.cancel_check():
raise GenerationCancelled()

def _get_gpt_response(self, messages, n=1):
response = []
max_tries = n + 2
n_tries = 0
while len(response) < n:
self._check_cancel()
s_time = time.time()
try:
print(f'\n\n{messages}\n\n')
Expand All @@ -57,20 +79,23 @@ def _get_gpt_response(self, messages, n=1):
top_p=self.top_p,
seed=self.seed,
stream=False,
max_completion_tokens=self.max_completion_tokens,
max_tokens=self.max_completion_tokens,
n=n,
)
except Exception as e:
self._check_cancel()
print(f'\nError: {e}\n\n')
n_tries += 1
if n_tries > max_tries:
each_response = '```\n\n[ERROR] Failed to generate\n\n```'
fallback = '```\n[ERROR] Failed to generate due to API error or quota.\n```'
response.append(fallback)
break
continue

print(f'\nTime consuming for one generation: {time.time()-s_time:.2f} seconds\n\n\n')

response.append(each_response.choices[0].message.content)
self._check_cancel()

if n == 1:
response = response[0]
Expand All @@ -82,6 +107,7 @@ def _get_gpt_o1_mini_response(self, messages, n=1):
max_tries = n + 2
n_tries = 0
while len(response) < n:
self._check_cancel()
s_time = time.time()
try:
print(f'\n\n{messages}\n\n')
Expand All @@ -91,13 +117,15 @@ def _get_gpt_o1_mini_response(self, messages, n=1):
temperature=self.temp,
seed=self.seed,
stream=False,
max_completion_tokens=self.max_completion_tokens,
max_tokens=self.max_completion_tokens,
n=n,
)
except Exception as e:
self._check_cancel()
print(f'\nError: {e}\n\n')
if "无可用渠道" in str(e):
time.sleep(2)
self._check_cancel()
continue

if "potentially violating our usage policy" in str(e) or 'bad response status' in str(e): # triggered by o1-mini
Expand All @@ -116,21 +144,25 @@ def _get_gpt_o1_mini_response(self, messages, n=1):
part_2_1 = '\n'.join(part_2_1_lines)
messages[1]['content'] = part_1 + '(with some details omitted):\n```\n' + part_2_1 + '\n```' + part_2_2

self._check_cancel()
continue

if "quota is not enough" in str(e):
time.sleep(10)
self._check_cancel()
continue

n_tries += 1
if n_tries > max_tries:
each_response = '```\n\n[ERROR] Failed to generate\n\n```'
fallback = '```\n[ERROR] Failed to generate due to API error or quota.\n```'
response.append(fallback)
break
continue

print(f'\nTime consuming for one generation: {time.time()-s_time:.2f} seconds\n\n\n')

response.append(each_response.choices[0].message.content)
self._check_cancel()

if n == 1:
response = response[0]
Expand All @@ -146,6 +178,7 @@ def _get_deepseek_qwen_response(self, messages, n=1, skip_deepseek_think: bool=F
messages[0]['content'] += '\n\n<think>\nSkip Thinking\n</think>\n\n'

while len(response) < n:
self._check_cancel()
s_time = time.time()
try:
each_response_raw = self.client.chat.completions.create(
Expand All @@ -154,10 +187,11 @@ def _get_deepseek_qwen_response(self, messages, n=1, skip_deepseek_think: bool=F
temperature=0.6,
seed=self.seed,
stream=False,
max_completion_tokens=self.max_completion_tokens,
max_tokens=self.max_completion_tokens,
n=1
)
except Exception as e:
self._check_cancel()
# the input is too long
if 'Please reduce the length' in str(e):
context_part = messages[0]['content'].split('(with some details omitted):')[1]
Expand All @@ -173,6 +207,11 @@ def _get_deepseek_qwen_response(self, messages, n=1, skip_deepseek_think: bool=F
messages[0]['content'] = messages[0]['content'].replace(context_part, reduced_context_part)

continue
# when API/quota errors persist, append fallback to avoid empty response
n_tries += 1
if n_tries >= max_tries:
response.append('```\n[ERROR] Failed to generate due to API error or quota.\n```')
break

print(f'Time consuming for one generation: {time.time()-s_time:.2f} seconds\n\n')
print(f'[INFO] Response:\n{each_response_raw.choices[0].message.content}\n\n\n')
Expand All @@ -191,6 +230,7 @@ def _get_deepseek_qwen_response(self, messages, n=1, skip_deepseek_think: bool=F
each_response = '```\nFailed to generate\n```'

response.append(each_response)
self._check_cancel()

if n == 1:
response = response[0]
Expand Down Expand Up @@ -395,4 +435,4 @@ def construct_prompt(self, gen_test_case, error_msg, target_focal_method, target

instruction += f"""# Output Requirements\nYour final output must strictly adhere to the following format:\n1: Begin with the exact prefix: "{self.gen_prefix}".\n2: End with the exact suffix: "{self.gen_suffix}".\nEnsure that no additional text appears before the prefix or after the suffix."""

return instruction
return instruction
6 changes: 4 additions & 2 deletions backend/configs.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,8 @@ def __init__(self, project_name, tester_path = '') -> None:
os.environ['OPENAI_BASE_URL'] = self.openai_url

self.project_name = project_name
self.llm_name = 'gpt-4o'
# allow overriding model via config.ini -> [openai] model = ...
self.llm_name = global_config['openai'].get('model', 'gpt-4o')

self.max_context_len = 1024
self.max_input_len = 4096
Expand All @@ -22,7 +23,8 @@ def __init__(self, project_name, tester_path = '') -> None:
else:
self.workspace = f'{self.root_dir}/intention_test_extension'

self.corpus_path = f'{self.workspace}/data/collected_coverages/{project_name}.json'
# Align with dump_collect_pairs: it writes to backend/data/{project}.json
self.corpus_path = f'{self.root_dir}/data/{project_name}.json'
self.project_without_test_file_path = f'{self.workspace}/data/repos_removing_test/{project_name}'
self.project_with_test_file_path = f'{self.workspace}/data/repos_with_test/{project_name}'

Expand Down
30 changes: 28 additions & 2 deletions backend/generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,8 @@

from pyexpat.errors import messages

from server import ModelQuerySession
from modules.session import ModelQuerySession
from modules.exceptions import GenerationCancelled
from configs import Configs
from agents import TestGenAgent, TestRefineAgent
from test_case_runner import TestCaseRunner
Expand All @@ -20,25 +21,36 @@ def __init__(self, configs: Configs, max_round=3, skip_deepseek_think: bool = Fa
self.test_refine_agent = TestRefineAgent(configs.llm_name, configs.project_name, configs.project_url, n_responses=1, skip_deepseek_think=skip_deepseek_think)
self.test_runner = TestCaseRunner(configs, configs.test_case_run_log_dir)
self.generation_with_refine_log = [] # [(test_status, prompt, test_case)]
self.query_session: ModelQuerySession | None = None
self._cancel_check = lambda: False
self._apply_cancel_hook()

def connect_to_request_session(self, query_session: ModelQuerySession):
self.query_session = query_session
self._apply_cancel_hook()

def update_messages_to_remote(self, messages):
# TODO notify front-end for messages, maybe trasmit full (instead of transmit update only)?
if self.query_session:
self.query_session.update_messages(messages)

def _ensure_not_cancelled(self):
if self.query_session and self.query_session.should_stop():
raise GenerationCancelled()

def generate_test_case_with_refine(self,
target_focal_method, target_context, target_test_case_desc, target_test_case_path,
referable_test_case, facts, junit_version,
prohibit_fact: bool = False, query_session: ModelQuerySession | None = None):
self.generation_with_refine_log = []
self.query_session = query_session
self._apply_cancel_hook()
self._ensure_not_cancelled()

target_test_class_name = target_test_case_path.split('/')[-1].replace('.java', '')
gen_test_case, prompt, messages = self.generate_test_case(target_focal_method, target_context, target_test_class_name, target_test_case_desc, referable_test_case, facts, junit_version, prohibit_fact)
self.update_messages_to_remote(messages)
self._ensure_not_cancelled()
error_msg, test_status = self.run_test_case(gen_test_case, target_test_case_path)
self.generation_with_refine_log.append((test_status, prompt, gen_test_case))

Expand All @@ -47,9 +59,11 @@ def generate_test_case_with_refine(self,
return gen_test_case, test_status, messages

for round in range(self.max_round):
self._ensure_not_cancelled()
gen_test_case, prompt, refine_messages = self.refine(gen_test_case, error_msg, target_focal_method, target_context, target_test_case_desc, target_test_case_path, facts, prohibit_fact)
messages += refine_messages
self.update_messages_to_remote(messages)
self._ensure_not_cancelled()
error_msg, test_status = self.run_test_case(gen_test_case, target_test_case_path)
self.generation_with_refine_log.append((test_status, prompt, gen_test_case))

Expand All @@ -61,21 +75,25 @@ def generate_test_case_with_refine(self,
return gen_test_case, test_status, messages

def finish_generate(self):
self._ensure_not_cancelled()
messages = self.test_gen_agent.generate_finish()
return messages

def generate_test_case(self, target_focal_method, target_context, target_test_class_name, target_test_case_desc, referable_test_case, facts, junit_version, prohibit_fact):
self._ensure_not_cancelled()
gen_test_case, prompt, messages = self.test_gen_agent.generate_test_case(target_focal_method, target_context, target_test_class_name, target_test_case_desc, referable_test_case, facts, junit_version, prohibit_fact)
return gen_test_case, prompt, messages

def refine(self, gen_test_case, error_msg, target_focal_method, target_context, target_test_case_desc, target_test_case_path, facts: list, prohibit_fact):
self._ensure_not_cancelled()
error_msg_lines = error_msg.split('\n')
error_msg_cut = '\n'.join(error_msg_lines[:self.max_line_error_msg])

refined_tc, prompt, messages = self.test_refine_agent.refine(gen_test_case, error_msg_cut, target_focal_method, target_context, target_test_case_desc, facts, prohibit_fact)
return refined_tc, prompt, messages

def run_test_case(self, test_case, test_case_path):
self._ensure_not_cancelled()
def _extract_error_msg(log):
error_msg = []
stop_flag = False
Expand Down Expand Up @@ -130,4 +148,12 @@ def _extract_error_msg(log):
error_msg = ""
test_status = 'success'

return error_msg, test_status
return error_msg, test_status

def _apply_cancel_hook(self):
def cancel_check() -> bool:
return bool(self.query_session and self.query_session.should_stop())

self._cancel_check = cancel_check
self.test_gen_agent.set_cancel_check(cancel_check)
self.test_refine_agent.set_cancel_check(cancel_check)
65 changes: 53 additions & 12 deletions backend/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,8 @@
from generator import IntentionTester
from dataset import Dataset
from configs import Configs
from server import ModelQuerySession
from modules.session import ModelQuerySession
from typing import Optional
import pathlib
from extension_api.collect_pairs.main import dump_collect_pairs

Expand Down Expand Up @@ -36,22 +37,47 @@ def load_corpus(self):
corpus_fm, corpus_fm_name, corpus_context, corpus_tc_name, corpus_test_case_path = [], [], [], [], []

for each_data in all_data:
corpus_fm.append(''.join(each_data['target_coverage']).replace('<COVER>', ''))
corpus_fm_name.append(each_data['focal_method_name'])
corpus_context.append(each_data['target_context'])
corpus_tc_name.append(each_data['target_test_case_name'].split('::::')[-1].split('(')[0])
corpus_test_case_path.append(each_data['focal_file_path'].replace('src/main/java', 'src/test/java').replace('.java', 'Test.java'))
if 'target_coverage' in each_data:
# original expected format
corpus_fm.append(''.join(each_data['target_coverage']).replace('<COVER>', ''))
corpus_fm_name.append(each_data.get('focal_method_name', ''))
corpus_context.append(each_data.get('target_context', ''))
tc_name = each_data.get('target_test_case_name', '')
corpus_tc_name.append(tc_name.split('::::')[-1].split('(')[0] if tc_name else '')
focal_file_path = each_data.get('focal_file_path', '')
if focal_file_path:
corpus_test_case_path.append(focal_file_path.replace('src/main/java', 'src/test/java').replace('.java', 'Test.java'))
else:
corpus_test_case_path.append('')
else:
# fallback to collect_pairs schema
# focal method text
fm = each_data.get('focal_method', [])
corpus_fm.append(''.join(fm) if isinstance(fm, list) else str(fm))
# focal method name
corpus_fm_name.append(each_data.get('focal_method_name', ''))
# use focal method content as context (best available without re-reading files)
corpus_context.append(''.join(fm) if isinstance(fm, list) else str(fm))
# derive test case simple name from test_name or test_path
test_name = each_data.get('test_name', '')
if test_name:
corpus_tc_name.append(test_name.split('(')[0].split('::::')[-1])
else:
test_path = each_data.get('test_path', '')
corpus_tc_name.append(os.path.splitext(os.path.basename(test_path))[0] if test_path else '')
# test case path provided directly
corpus_test_case_path.append(each_data.get('test_path', ''))

self.corpus = {
'corpus_fm': corpus_fm,
'corpus_fm_name': corpus_fm_name,
'corpus_fm': corpus_fm,
'corpus_fm_name': corpus_fm_name,
'corpus_context': corpus_context,
'corpus_tc_name': corpus_tc_name,
'corpus_tc_name': corpus_tc_name,
'corpus_test_case_path': corpus_test_case_path
}
}


def main(target_focal_method, target_focal_file, test_desc, project_path, focal_file_path, query_session: ModelQuerySession | None = None):
def main(target_focal_method, target_focal_file, test_desc, project_path, focal_file_path, query_session: Optional[ModelQuerySession] = None):
# project_name = project_path.split('/')[-1] not compatible with Windows path
project_name = pathlib.Path(project_path).stem
# replace the disk letter to upper case to match CodeQL path
Expand Down Expand Up @@ -98,7 +124,22 @@ def main(target_focal_method, target_focal_file, test_desc, project_path, focal_
target_test_case_desc = test_desc_data['test_desc']['under_setting']

# TODO LSP now cannot run in Windows
offline_fact_ref_data = dataset.load_offline_fact_ref_data()
try:
offline_fact_ref_data = dataset.load_offline_fact_ref_data()
except FileNotFoundError:
# Fallback: construct empty facts/references with proper length
corpus_len = len(intention_test.corpus['corpus_fm_name']) if intention_test.corpus else 0
offline_fact_ref_data = [
{
'target_coverage_idx': i,
'rag_references': [],
'disc_facts': [],
'disc_facts_sim': [],
'top_usages': [],
'top_usages_sim': []
}
for i in range(corpus_len)
]

# prepare test generator
dtester = IntentionTester(configs)
Expand Down
5 changes: 5 additions & 0 deletions backend/modules/exceptions.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
class GenerationCancelled(Exception):
"""Raised when a generation session is cancelled by the user."""

def __init__(self, message: str = "Generation cancelled by user") -> None:
super().__init__(message)
Loading