diff --git a/backend/agents.py b/backend/agents.py index d6b801f..6591427 100644 --- a/backend/agents.py +++ b/backend/agents.py @@ -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): @@ -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) @@ -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') @@ -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] @@ -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') @@ -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 @@ -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] @@ -146,6 +178,7 @@ def _get_deepseek_qwen_response(self, messages, n=1, skip_deepseek_think: bool=F messages[0]['content'] += '\n\n\nSkip Thinking\n\n\n' while len(response) < n: + self._check_cancel() s_time = time.time() try: each_response_raw = self.client.chat.completions.create( @@ -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] @@ -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') @@ -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] @@ -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 \ No newline at end of file + return instruction diff --git a/backend/configs.py b/backend/configs.py index a433a58..65ff9c4 100644 --- a/backend/configs.py +++ b/backend/configs.py @@ -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 @@ -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}' diff --git a/backend/generator.py b/backend/generator.py index 71f0a9c..bbae9fc 100644 --- a/backend/generator.py +++ b/backend/generator.py @@ -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 @@ -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)) @@ -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)) @@ -61,14 +75,17 @@ 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]) @@ -76,6 +93,7 @@ def refine(self, gen_test_case, error_msg, target_focal_method, target_context, 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 @@ -130,4 +148,12 @@ def _extract_error_msg(log): error_msg = "" test_status = 'success' - return error_msg, test_status \ No newline at end of file + 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) diff --git a/backend/main.py b/backend/main.py index 4523c31..9437c2a 100644 --- a/backend/main.py +++ b/backend/main.py @@ -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 @@ -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('', '')) - 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('', '')) + 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 @@ -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) diff --git a/backend/modules/exceptions.py b/backend/modules/exceptions.py new file mode 100644 index 0000000..06b6287 --- /dev/null +++ b/backend/modules/exceptions.py @@ -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) diff --git a/backend/modules/messages.py b/backend/modules/messages.py new file mode 100644 index 0000000..195c262 --- /dev/null +++ b/backend/modules/messages.py @@ -0,0 +1,42 @@ +from __future__ import annotations + +import json +from dataclasses import dataclass +from typing import Any, Dict, Union + + +def _to_bytes(payload: Dict[str, Any]) -> bytes: + return json.dumps(payload).encode() + + +@dataclass +class StatusMessage: + status: str + message: Union[str, Dict[str, Any]] = "" + + def to_bytes(self) -> bytes: + return _to_bytes( + { + "type": "status", + "data": { + "status": self.status, + "message": self.message, + }, + } + ) + + +@dataclass +class ModelMessage: + data: Dict[str, Any] + + def to_bytes(self) -> bytes: + return _to_bytes({"type": "msg", "data": self.data}) + + +@dataclass +class NoRefMessage: + data: Dict[str, Any] + + def to_bytes(self) -> bytes: + return _to_bytes({"type": "noreference", "data": self.data}) diff --git a/backend/modules/registry.py b/backend/modules/registry.py new file mode 100644 index 0000000..1584a34 --- /dev/null +++ b/backend/modules/registry.py @@ -0,0 +1,30 @@ +from __future__ import annotations + +import threading +from typing import Dict, Iterable, Optional + +from .session import ModelQuerySession + + +class SessionRegistry: + """线程安全的会话注册表。""" + + def __init__(self) -> None: + self._sessions: Dict[str, ModelQuerySession] = {} + self._lock = threading.Lock() + + def register(self, session: ModelQuerySession) -> None: + with self._lock: + self._sessions[session.session_id] = session + + def remove(self, session_id: str) -> None: + with self._lock: + self._sessions.pop(session_id, None) + + def get(self, session_id: str) -> Optional[ModelQuerySession]: + with self._lock: + return self._sessions.get(session_id) + + def list_active_ids(self) -> Iterable[str]: + with self._lock: + return tuple(self._sessions.keys()) diff --git a/backend/modules/session.py b/backend/modules/session.py new file mode 100644 index 0000000..a2f952b --- /dev/null +++ b/backend/modules/session.py @@ -0,0 +1,85 @@ +from __future__ import annotations + +import logging +import threading +from typing import Any, Callable, Dict, List + +from .exceptions import GenerationCancelled +from .messages import ModelMessage, NoRefMessage, StatusMessage + +logger = logging.getLogger(__name__) + +ResponseWriter = Callable[[bytes], None] +QueryExecutor = Callable[[Dict[str, Any], "ModelQuerySession"], None] + + +class ModelQuerySession: + """封装单次生成流程的上下文与与客户端通信能力。""" + + required_fields = [ + "target_focal_method", + "target_focal_file", + "test_desc", + "project_path", + "focal_file_path", + ] + + def __init__( + self, + session_id: str, + raw_data: Dict[str, Any], + writer: ResponseWriter, + executor: QueryExecutor, + junit_version: int, + ) -> None: + self.session_id = session_id + self.raw_data = raw_data + self._writer = writer + self._executor = executor + self.junit_version = junit_version + + self.messages: List[Dict[str, Any]] = [] + self.query_data = {field: self.raw_data[field] for field in self.required_fields} + self._session_running = False + self._cancel_event = threading.Event() + + def start_query(self) -> None: + if self._session_running: + logger.warning("Session %s already running", self.session_id) + return + self._session_running = True + logger.info("Starting query session %s", self.session_id) + try: + self._executor(self.query_data, self) + except GenerationCancelled: + logger.info("Query session %s cancelled by user", self.session_id) + finally: + self._session_running = False + + def update_messages(self, messages: List[Dict[str, Any]]) -> None: + self.messages = messages + data_to_send = {"session_id": self.session_id, "messages": messages} + self._safe_write(ModelMessage(data_to_send).to_bytes()) + + def write_start_message(self) -> None: + self._safe_write(StatusMessage("start", {"session_id": self.session_id}).to_bytes()) + + def write_noref_message(self) -> None: + payload = {"session_id": self.session_id, "junit_version": self.junit_version} + self._safe_write(NoRefMessage(payload).to_bytes()) + + def write_finish_message(self) -> None: + self._safe_write(StatusMessage("finish", {"session_id": self.session_id}).to_bytes()) + + def request_stop(self) -> None: + self._cancel_event.set() + + def should_stop(self) -> bool: + return self._cancel_event.is_set() + + def _safe_write(self, payload: bytes) -> None: + try: + self._writer(payload) + except BrokenPipeError: + logger.warning("Connection closed for session %s", self.session_id) + self.request_stop() diff --git a/backend/server.py b/backend/server.py index e7f7774..384cabf 100644 --- a/backend/server.py +++ b/backend/server.py @@ -1,240 +1,194 @@ -import socket -import http.server -import socketserver -import threading -import datetime +from __future__ import annotations + +import argparse import json -from time import strftime -from xml.etree.ElementPath import prepare_child import logging -import sys -import traceback -import argparse -import main -import hashlib +import socketserver +import threading +from http.server import BaseHTTPRequestHandler +from typing import Any, Dict -port = 8080 +try: + from backend import main as generation_entry_module # when run as package +except ImportError: + import main as generation_entry_module # when invoked from backend directory +from modules.registry import SessionRegistry +from modules.session import ModelQuerySession -# a standard python logger logger = logging.getLogger(__name__) -# basiConfig can only be called once -logging.basicConfig(stream=sys.stdout, level=logging.INFO, format='[%(asctime)s] [%(levelname)s] %(message)s') - -global_junit_version = 4 - -class StatusMessage: - def __init__(self, status: str, message: str | dict = ''): - self.status = status - self.message = message - - def response(self): - return json.dumps({ - "type": "status", - "data": { - "status": self.status, - "message": self.message - } - }).encode() - -class ModelMessage: - def __init__(self, data): - self.data = data - - def response(self): - return json.dumps({ - "type": "msg", - "data": self.data - }).encode() - -class NoRefMessage: - def __init__(self, data): - self.data = data - - def response(self): - return json.dumps({ - "type": "noreference", - "data": self.data - }).encode() - -class QueryHandler(http.server.BaseHTTPRequestHandler): - def do_POST(self): - global global_junit_version - - if self.path == '/session': - - try: - self.request.settimeout(2.0) - query_text_bytes = self.rfile.read(int(self.headers['Content-Length'])) - query_text = query_text_bytes.decode('utf-8') - self.request.settimeout(None) - - try: - query_session = assign_to_session(query_text, self) - except Exception as e: - logger.error(f'Request may be invalid. Message: {e}. Request:\n{query_text}\n{traceback.format_exc()}') - self.end_with_request_error(str(e)) - return - - if query_session: - self.send_keep_alive_header() - # self.write_single_line(StatusMessage('start').response()) - query_session.write_start_message() - query_session.start_query() - # self.write_single_line(StatusMessage('finish').response()) - query_session.write_finish_message() - # no need to flush because the handle_one_request will do that - self.end_session() - else: - raise ValueError("No query session can be constructed or retrieved from request") - - except Exception as e: - logger.error(f'Error handling request. Message: {e}. Request:\n{self.request}\n{traceback.format_exc()}') - self.end_with_internal_error(str(e)) - - elif self.path == '/junitVersion': - try: - self.request.settimeout(2.0) - junit_version = int(json.loads(self.rfile.read(int(self.headers['Content-Length'])).decode('utf-8'))['data']) - self.request.settimeout(None) - - global_junit_version = junit_version - except Exception as e: - logger.error(f"Error handling request. Message: {e}. Request:\n{self.request}\n{traceback.format_exc()}") - self.end_with_internal_error(str(e)) +logging.basicConfig( + level=logging.INFO, + format="[%(asctime)s] [%(levelname)s] %(message)s", +) + +DEFAULT_PORT = 8080 +_global_junit_version = 4 +_session_registry = SessionRegistry() + + +class ThreadedTCPServer(socketserver.ThreadingMixIn, socketserver.TCPServer): + daemon_threads = True + allow_reuse_address = True + + +class ResponseStream: + """封装 Handler 的写操作,确保线程安全。""" + + def __init__(self, handler: BaseHTTPRequestHandler) -> None: + self._handler = handler + self._lock = threading.Lock() + + def __call__(self, data: bytes) -> None: + with self._lock: + self._handler.wfile.write(data + b"\n") + self._handler.wfile.flush() + + +def run_generation(query_data: Dict[str, Any], session: ModelQuerySession) -> None: + generation_entry_module.main(**query_data, query_session=session) + + +def build_session(payload: Dict[str, Any], handler: BaseHTTPRequestHandler) -> ModelQuerySession: + session_id = payload["session_id"] + response_stream = ResponseStream(handler) + return ModelQuerySession( + session_id=session_id, + raw_data=payload["data"], + writer=response_stream, + executor=run_generation, + junit_version=_global_junit_version, + ) + +def validate_query_payload(payload: Dict[str, Any]) -> Dict[str, Any]: + if payload.get("type") != "query": + raise ValueError("Unsupported request type") + data = payload.get("data") + if not isinstance(data, dict): + raise ValueError("Query data must be a JSON object") + missing = [field for field in ModelQuerySession.required_fields if field not in data] + if missing: + raise ValueError(f"Missing required fields: {', '.join(missing)}") + return {"session_id": payload.get("session_id") or payload.get("id") or handler_uuid(), "data": data} + + +def handler_uuid() -> str: + import uuid + + return uuid.uuid4().hex + + +class QueryHandler(BaseHTTPRequestHandler): + server_version = "IntentionTestHTTP/1.0" + + def do_POST(self) -> None: # noqa: N802 + if self.path == "/session": + self._handle_session_request() + elif self.path == "/session/stop": + self._handle_stop_request() + elif self.path == "/junitVersion": + self._handle_junit_version() else: self.send_response(404) self.end_headers() - def send_keep_alive_header(self): - self.send_response(200, 'Success') - self.send_header('Content-type', 'application/json') # this doesn't exist on - self.send_header('Cache-Control', 'no-cache') - self.send_header('Connection', 'keep-alive') + def _handle_session_request(self) -> None: + try: + payload = self._read_json_body() + request_payload = validate_query_payload(payload) + session = build_session(request_payload, self) + except Exception as exc: # broad catch to surface payload issues + logger.error("Invalid session request: %s", exc, exc_info=True) + self._end_with_error(400, "Bad Request", str(exc)) + return + + try: + _session_registry.register(session) + self._send_keep_alive_header() + session.write_start_message() + session.start_query() + session.write_finish_message() + except Exception as exc: + logger.error("Error processing session: %s", exc, exc_info=True) + self._end_with_error(500, "Internal Server Error", str(exc)) + finally: + _session_registry.remove(session.session_id) + self._end_session() + + def _handle_stop_request(self) -> None: + try: + payload = self._read_json_body() + session_id = payload.get("session_id") + if not session_id: + raise ValueError("Missing session_id") + session = _session_registry.get(session_id) + if not session: + self.send_response(404, "Session Not Found") + self.end_headers() + return + session.request_stop() + self.send_response(200, "Stopping") + self.end_headers() + except ValueError as exc: + self._end_with_error(400, "Bad Request", str(exc)) + except Exception as exc: + logger.error("Failed to stop session: %s", exc, exc_info=True) + self._end_with_error(500, "Internal Server Error", str(exc)) + + def _handle_junit_version(self) -> None: + global _global_junit_version + + try: + payload = self._read_json_body() + version = int(payload["data"]) + except Exception as exc: + self._end_with_error(400, "Bad Request", f"Invalid payload: {exc}") + return + + _global_junit_version = version + self.send_response(200, "Success") self.end_headers() - - def end_with_error(self, code: int, error_msg: str, concrete_msg: str): + + def _send_keep_alive_header(self) -> None: + self.send_response(200, "Success") + self.send_header("Content-type", "application/json") + self.send_header("Cache-Control", "no-cache") + self.send_header("Connection", "keep-alive") + self.end_headers() + + def _end_with_error(self, code: int, error_msg: str, _: str) -> None: self.send_response(code, error_msg) self.end_headers() - self.end_session() - # self.wfile.write(StatusMessage('error', concrete_msg).response()) + self._end_session() - def end_with_request_error(self, msg: str): - self.end_with_error(400, 'Bad Request', msg) + def _end_session(self) -> None: + self.close_connection = True - def end_with_internal_error(self, msg: str): - self.end_with_error(500, 'Internal Server Error', msg) + def _read_json_body(self) -> Dict[str, Any]: + content_length = int(self.headers.get("Content-Length", 0)) + body = self.rfile.read(content_length).decode("utf-8") + return json.loads(body) if body else {} - def end_session(self): - self.close_connection = True - def write_single_line(self, data: bytes): - self.wfile.write(data + b'\n') - self.wfile.flush() - -class ModelQuerySession: - '''Session persistent data.''' - required_fields = ['target_focal_method', 'target_focal_file', 'test_desc', 'project_path', 'focal_file_path'] - - def __init__(self, session_id: str, raw_data: dict, handler: QueryHandler): - self.session_id = session_id - self.raw_data = raw_data - self.handler = handler - self.messages = [] - self.junit_version = global_junit_version - - self.query_data = self.prepare_query_arguments() - self.session_running = False - - def prepare_query_arguments(self): - # do with session_meta_data - return {x: self.raw_data[x] for x in self.required_fields } - - def start_query(self): - if not self.session_running: - self.session_running = True - logger.info(f'Starting query session {self.session_id}') - main.main(**self.query_data, query_session = self) - self.session_running = False - - def update_messages(self, messages): - self.messages = messages - data_to_send = { - 'session_id': self.session_id, - 'messages': messages - } - self.handler.write_single_line(ModelMessage(data_to_send).response()) - - def write_start_message(self): - data = { - 'session_id': self.session_id - } - self.handler.write_single_line(StatusMessage('start', data).response()) - - def write_noref_message(self): - data = { - 'session_id': self.session_id, - 'junit_version': self.junit_version - } - self.handler.write_single_line(NoRefMessage(data).response()) - - def write_finish_message(self): - data = { - 'session_id': self.session_id - } - self.handler.write_single_line(StatusMessage('finish', data).response()) - -# Not used now, we still send raw time -def get_hash(s: str): - h = hashlib.sha256(s.encode('utf-8')) - return h.hexdigest() - -sessions: dict[str, ModelQuerySession] = {} - -def assign_to_session(query_text: str, query_handler: QueryHandler) -> ModelQuerySession | None: - # do with sessions - query_data = json.loads(query_text) - if query_data['type'] != 'query': - raise NotImplementedError('None query is not supported yet') - - time_str = datetime.datetime.now(datetime.timezone.utc).strftime('%Y-%m-%d %H:%M:%S %Z') - new_session = ModelQuerySession(time_str, query_data['data'], query_handler) - return new_session - # TODO sometimes session should be retrived, return None if not found - -# def find_open_port(): -# with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: -# s.bind(('', 0)) # Bind to any available port -# return s.getsockname()[1] # Return the port number - -def start_http_server(port: int): - logger.info(f'Starting HTTP server on port {port}') - - httpd = socketserver.TCPServer(("", port), QueryHandler) - port = httpd.server_address[1] - logger.info(f'HTTP server is started and listening on {port}') - - th = threading.Thread(target=httpd.serve_forever, daemon=True) - th.start() - th.join() - -# class StdioServer: -# def __init__(self): -# self.handler = QueryHandler() -# self.should_run = True - -# def serve_forever(self): -# while self.should_run: -# request = sys.stdin.read() -# if request: -# self.handler.handle_one_request() - -# def shutdown(self): -# self.should_run = False - -if __name__ == '__main__': - parser = argparse.ArgumentParser(description='Start the model server') - parser.add_argument('--port', type=int, default=8080, help='Port to start the server on') # by default listen to a random port - +def start_http_server(port: int) -> None: + logger.info("Starting HTTP server on port %s", port) + httpd = ThreadedTCPServer(("", port), QueryHandler) + actual_port = httpd.server_address[1] + logger.info("HTTP server is listening on %s", actual_port) + try: + httpd.serve_forever() + except KeyboardInterrupt: + logger.info("Shutting down HTTP server") + finally: + httpd.server_close() + + +def main() -> None: + parser = argparse.ArgumentParser(description="Start the model server") + parser.add_argument("--port", type=int, default=DEFAULT_PORT, help="Port to start the server on") args = parser.parse_args() start_http_server(args.port) + + +if __name__ == "__main__": + main() diff --git a/src/client.ts b/src/client.ts index c319a8c..1e091fa 100644 --- a/src/client.ts +++ b/src/client.ts @@ -1,11 +1,15 @@ // create a python subprocess and communicate with it through network -import { request, RequestOptions } from 'http'; +import { request, RequestOptions, ClientRequest } from 'http'; export class TesterSession { private updateMessageCallback?: (...args: any[]) => any; private errorCallbcak?: (...args: any[]) => any; private showNoRefMsg?: (...args: any[]) => any; private connectToPort: number; + private currentRequest?: ClientRequest; + private finishActiveRequest?: () => void; + private isCancelling = false; + private activeSessionId?: string; // setting connectToPort to 0 to start up an internal server constructor(updateMessageCallback?: (...args: any[]) => any, errorCallback?: (...args: any[]) => any, showNoRefMsg?: (...args: any[]) => any, connectToPort: number = 0) { @@ -53,6 +57,7 @@ export class TesterSession { async startQuery(args: any, cancelCb: (e: any) => any) { const requestData = new TextEncoder().encode(JSON.stringify({ type: 'query', data: args }) + '\n'); + this.activeSessionId = undefined; const options: RequestOptions = { hostname: 'localhost', @@ -66,7 +71,12 @@ export class TesterSession { }; let finish: (value?: any) => void; - const finishePromise = new Promise((res, rej) => { finish = res; }); + const finishePromise = new Promise((res) => { finish = res; }); + this.finishActiveRequest = () => { + finish(); + this.resetRequestState(); + }; + this.isCancelling = false; const req = request(options, (res) => { let status = 'before-start'; @@ -82,13 +92,15 @@ export class TesterSession { if (!(msg.type && msg.data && msg.type === 'status' && msg.data.status === 'start')) { throw TypeError('Failed to receive start message'); } + this.activeSessionId = msg.data.session_id; status = 'started'; } else if (status !== 'finished') { // receive messages if (msg.type && msg.data) { if (msg.type === 'status' && msg.data.status === 'finish') { status = 'finished'; - finish(); + this.activeSessionId = undefined; + this.finishActiveRequest?.(); return; } else if (msg.type === 'msg' && msg.data.session_id && msg.data.messages) { if (this.updateMessageCallback) { @@ -109,26 +121,88 @@ export class TesterSession { } } catch (e) { - console.error(e); - cancelCb(e); + if (!this.isCancelling) { + console.error(e); + cancelCb(e); + } } }); res.on('end', () => { console.log('No more data in response.'); - // this.close(); + if (!this.isCancelling) { + this.resetRequestState(); + } }); res.on('error', (e) => { + if (this.isCancelling) { + return; + } console.error(e); }); }); + this.currentRequest = req; req.on('error', (e) => { + if (this.isCancelling) { + return; + } console.error(`Problem on request: ${e}`); }); req.write(requestData); req.end(); await finishePromise; + this.resetRequestState(); + } + + public cancelCurrentQuery(): void { + if (!this.currentRequest) { + return; + } + this.isCancelling = true; + this.currentRequest.destroy(); + this.finishActiveRequest?.(); + } + + public async stopActiveSession(): Promise { + this.cancelCurrentQuery(); + await this.sendStopSignal(); + } + + private resetRequestState(): void { + this.currentRequest = undefined; + this.finishActiveRequest = undefined; + this.isCancelling = false; + } + + private async sendStopSignal(): Promise { + if (!this.activeSessionId) { + return; + } + const payload = Buffer.from(JSON.stringify({ session_id: this.activeSessionId }), 'utf-8'); + const options: RequestOptions = { + hostname: 'localhost', + port: this.connectToPort, + path: '/session/stop', + method: 'POST', + headers: { + 'Content-Type': 'application/json', + 'Content-Length': payload.length.toString() + } + }; + + await new Promise((resolve) => { + const req = request(options, () => { + resolve(); + }); + req.on('error', (err) => { + console.error(`Failed to stop backend session: ${err}`); + resolve(); + }); + req.write(payload); + req.end(); + }); + this.activeSessionId = undefined; } } diff --git a/src/extension.ts b/src/extension.ts index 5c1a22d..035a046 100644 --- a/src/extension.ts +++ b/src/extension.ts @@ -8,9 +8,12 @@ import { marked } from 'marked'; import { ExtensionMetadata } from './constants'; import { showANewEditorForInput } from './utils'; +let activeSession: TesterSession | undefined; + export function activate(context: vscode.ExtensionContext): void { const viewId = 'testView.sidebar'; const testerWebViewProvider = new TesterWebViewProvider(context); + testerWebViewProvider.setMessageHandler((msg) => handleWebviewCommand(msg, testerWebViewProvider)); context.subscriptions.push( vscode.window.registerWebviewViewProvider(viewId, testerWebViewProvider, { @@ -132,15 +135,24 @@ async function generateTest(focalMethod: string, focalFile: string, testDesc: st }, connectToPort ); + activeSession = session; + await sendSessionState(ui, 'running'); await ui.showMessage({ role: 'system-wait', content: 'Server is preparing...' }); // await session.connect(); - await session.startQuery(generateParams, (e: any) => { - vscode.window.showErrorMessage(`Query error when connecting to the server: ${e}`); - // ui.showMessage({ cmd: 'error', message: 'an error has occurred'}); - }); + try { + await session.startQuery(generateParams, (e: any) => { + vscode.window.showErrorMessage(`Query error when connecting to the server: ${e}`); + // ui.showMessage({ cmd: 'error', message: 'an error has occurred'}); + }); + } finally { + if (activeSession === session) { + activeSession = undefined; + } + await sendSessionState(ui, 'idle'); + } } // TODO add blocking to prevent 2 sessions at the same time, or allow parallel sessions in new tab @@ -198,4 +210,34 @@ async function showWait(msg: any, ui: TesterWebViewProvider): Promise { export function deactivate() { } +async function handleWebviewCommand(msg: any, ui: TesterWebViewProvider): Promise { + if (!(msg && msg.cmd)) { + return; + } + if (msg.cmd === 'stop-run') { + await stopActiveSession(ui); + } else if (msg.cmd === 'clear-chat') { + await ui.showMessage({ cmd: 'clear', toIndex: 0 }); + } +} +async function stopActiveSession(ui: TesterWebViewProvider): Promise { + if (!activeSession) { + await sendSessionState(ui, 'idle'); + return; + } + try { + await activeSession.stopActiveSession(); + } finally { + activeSession = undefined; + } + await sendSessionState(ui, 'stopped', '生成已被手动停止。'); +} + +async function sendSessionState( + ui: TesterWebViewProvider, + state: 'idle' | 'running' | 'stopped', + message?: string +): Promise { + await ui.showMessage({ cmd: 'session-state', state, message }); +} diff --git a/src/sidebarView.ts b/src/sidebarView.ts index e25920c..a309587 100644 --- a/src/sidebarView.ts +++ b/src/sidebarView.ts @@ -1,7 +1,6 @@ import * as vscode from 'vscode'; import * as path from 'path'; import * as fs from 'fs'; -import * as marked from 'marked'; import { detectCodeLang, extractGenTestCode, extractRefTestCode, isGenTestPrompt, langSuffix, shouldGenTestPrompt } from './textUtils'; import { CodeHistoryDiffPlayer } from './diffView'; @@ -15,6 +14,7 @@ export function setWebRoot(root: string) { export class TesterWebViewProvider implements vscode.WebviewViewProvider { private _context: vscode.ExtensionContext; private _view?: vscode.Webview; + private _messageHandler?: (msg: any) => Thenable | void; constructor(context: vscode.ExtensionContext) { this._context = context; @@ -34,11 +34,17 @@ export class TesterWebViewProvider implements vscode.WebviewViewProvider { if (msg.cmd === 'open-code' && msg.content && msg.lang) { const doc = await vscode.workspace.openTextDocument({ language: msg.lang, content: msg.content }); vscode.window.showTextDocument(doc); + } else if (this._messageHandler) { + await this._messageHandler(msg); } }); } + public setMessageHandler(handler: (msg: any) => Thenable | void): void { + this._messageHandler = handler; + } + private getHtmlContent(): string { const htmlPath = path.join(webRoot, 'index.html'); return fs.readFileSync(htmlPath, 'utf8'); diff --git a/web/index.html b/web/index.html index 9ce6f10..8d84e42 100644 --- a/web/index.html +++ b/web/index.html @@ -6,171 +6,33 @@ Chat Bot - - + + + + Intention Test + LLM-based iterative test assistant + + + + Idle + + + 跳到最新 + 阅读锁定 + 停止 + 清空 + + Intention Test 🧪 An LLM-based iterative test generator. Waiting for request... - -