From 6c0a622a293e66cffad5a33da03db61a2cead8ba Mon Sep 17 00:00:00 2001
From: Lyican <1260147616@qq.com>
Date: Sat, 3 Jan 2026 18:15:00 +0800
Subject: [PATCH 1/4] cicd: ci test case
---
backend/tests/__init__.py | 1 +
backend/tests/conftest.py | 26 ++++
backend/tests/test_agents.py | 216 ++++++++++++++++++++++++++++++
backend/tests/test_core.py | 95 +++++++++++++
backend/tests/test_dataset.py | 67 +++++++++
backend/tests/test_integration.py | 67 +++++++++
backend/tests/test_messages.py | 71 ++++++++++
backend/tests/test_retriever.py | 133 ++++++++++++++++++
backend/tests/test_runner.py | 113 ++++++++++++++++
backend/tests/test_server.py | 146 ++++++++++++++++++++
10 files changed, 935 insertions(+)
create mode 100644 backend/tests/__init__.py
create mode 100644 backend/tests/conftest.py
create mode 100644 backend/tests/test_agents.py
create mode 100644 backend/tests/test_core.py
create mode 100644 backend/tests/test_dataset.py
create mode 100644 backend/tests/test_integration.py
create mode 100644 backend/tests/test_messages.py
create mode 100644 backend/tests/test_retriever.py
create mode 100644 backend/tests/test_runner.py
create mode 100644 backend/tests/test_server.py
diff --git a/backend/tests/__init__.py b/backend/tests/__init__.py
new file mode 100644
index 0000000..09d41c4
--- /dev/null
+++ b/backend/tests/__init__.py
@@ -0,0 +1 @@
+# Backend tests package
diff --git a/backend/tests/conftest.py b/backend/tests/conftest.py
new file mode 100644
index 0000000..54b388a
--- /dev/null
+++ b/backend/tests/conftest.py
@@ -0,0 +1,26 @@
+"""
+Pytest configuration and fixtures for backend tests.
+"""
+import os
+import sys
+import pytest
+
+# Add backend directory to Python path for imports
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+
+
+@pytest.fixture(autouse=True)
+def setup_env():
+ """设置测试所需的环境变量,避免 OpenAI 客户端初始化失败"""
+ os.environ.setdefault('OPEN_AI_KEY', 'test-api-key')
+ os.environ.setdefault('OPENAI_BASE_URL', 'https://api.test.openai.com/v1')
+
+ # Ensure NLTK stopwords are available for retriever tests.
+ try:
+ import nltk
+ from nltk.corpus import stopwords
+
+ stopwords.words('english')
+ except LookupError:
+ nltk.download('stopwords', quiet=True)
+ yield
diff --git a/backend/tests/test_agents.py b/backend/tests/test_agents.py
new file mode 100644
index 0000000..3b4a134
--- /dev/null
+++ b/backend/tests/test_agents.py
@@ -0,0 +1,216 @@
+"""
+Tests for agents.py utility methods.
+
+These tests cover the pure utility functions in the Agent class
+that don't require external API calls.
+"""
+
+
+class TestAgentExtractCode:
+ """测试 Agent.extract_code_from_response 方法"""
+
+ def test_extract_java_code_block(self):
+ """从 ```java 代码块中提取代码"""
+ from agents import Agent
+
+ agent = Agent('gpt-4o')
+ response = '''Here is the code:
+```java
+public class Test {
+ public void test() {
+ System.out.println("Hello");
+ }
+}
+```
+'''
+ result = agent.extract_code_from_response(response)
+ assert 'public class Test' in result
+ assert 'System.out.println' in result
+
+ def test_extract_plain_code_block(self):
+ """从普通 ``` 代码块中提取代码"""
+ from agents import Agent
+
+ agent = Agent('gpt-4o')
+ response = '''Some text
+```
+def hello():
+ print("world")
+```
+'''
+ result = agent.extract_code_from_response(response)
+ assert 'def hello()' in result
+ assert 'print("world")' in result
+
+ def test_extract_no_code_block_returns_space(self):
+ """没有代码块时返回空格"""
+ from agents import Agent
+
+ agent = Agent('gpt-4o')
+ response = 'This is just plain text without any code block.'
+ result = agent.extract_code_from_response(response)
+ assert result == " "
+
+ def test_extract_multiple_code_blocks_takes_first(self):
+ """多个代码块时取第一个"""
+ from agents import Agent
+
+ agent = Agent('gpt-4o')
+ response = '''First block:
+```java
+class First {}
+```
+Second block:
+```java
+class Second {}
+```
+'''
+ result = agent.extract_code_from_response(response)
+ assert 'First' in result
+
+
+class TestAgentLineNumbers:
+ """测试 Agent 的行号处理方法"""
+
+ def test_add_line_numbers(self):
+ """测试添加行号"""
+ from agents import Agent
+
+ agent = Agent('gpt-4o')
+ content = "line1\nline2\nline3"
+ result = agent.add_line_numbers(content)
+
+ assert result == "1:line1\n2:line2\n3:line3"
+
+ def test_remove_line_numbers(self):
+ """测试移除行号"""
+ from agents import Agent
+
+ agent = Agent('gpt-4o')
+ content = "1:line1\n2:line2\n3:line3"
+ result = agent.remove_line_numbers(content)
+
+ assert result == "line1\nline2\nline3"
+
+ def test_add_remove_roundtrip(self):
+ """添加后移除应该还原原始内容"""
+ from agents import Agent
+
+ agent = Agent('gpt-4o')
+ original = "public class Test {\n void method() {}\n}"
+
+ with_numbers = agent.add_line_numbers(original)
+ restored = agent.remove_line_numbers(with_numbers)
+
+ assert restored == original
+
+
+class TestAgentRemoveThinking:
+ """测试 Agent.remove_thinking 方法"""
+
+ def test_remove_thinking_with_tag(self):
+ """有 标签时提取标签后的内容"""
+ from agents import Agent
+
+ agent = Agent('gpt-4o')
+ response = "\nSome thinking process\n\nActual answer here"
+ result = agent.remove_thinking(response)
+
+ assert result == "Actual answer here"
+
+ def test_remove_thinking_no_tag_returns_none(self):
+ """没有 标签时返回 None"""
+ from agents import Agent
+
+ agent = Agent('gpt-4o')
+ response = "Just a normal response without thinking tag"
+ result = agent.remove_thinking(response)
+
+ assert result is None
+
+ def test_remove_thinking_empty_after_tag(self):
+ """标签后为空时返回空字符串"""
+ from agents import Agent
+
+ agent = Agent('gpt-4o')
+ response = "\nThinking...\n\n "
+ result = agent.remove_thinking(response)
+
+ assert result == ""
+
+
+class TestTestDescAgentCheckGeneration:
+ """测试 TestDescAgent.check_generation 方法"""
+
+ def test_valid_format_returns_true(self):
+ """格式正确时返回 True"""
+ from agents import TestDescAgent
+
+ agent = TestDescAgent('gpt-4o')
+ desc = """# Objective
+Test something
+
+# Preconditions
+1. Some precondition
+
+# Expected Results
+1. Some result
+"""
+ assert agent.check_generation(desc) is True
+
+ def test_missing_section_returns_false(self):
+ """缺少部分时返回 False"""
+ from agents import TestDescAgent
+
+ agent = TestDescAgent('gpt-4o')
+ desc = """# Objective
+Test something
+
+# Preconditions
+1. Some precondition
+"""
+ assert agent.check_generation(desc) is False
+
+ def test_duplicate_section_returns_false(self):
+ """重复部分时返回 False"""
+ from agents import TestDescAgent
+
+ agent = TestDescAgent('gpt-4o')
+ desc = """# Objective
+Test something
+
+# Objective
+Another objective
+
+# Preconditions
+1. Some precondition
+
+# Expected Results
+1. Some result
+"""
+ assert agent.check_generation(desc) is False
+
+
+class TestAgentRemoveSingleLineNumber:
+ """Test Agent.remove_single_line_number method."""
+
+ def test_removes_line_number(self):
+ """Test removing a single line number."""
+ from agents import Agent
+
+ agent = Agent.__new__(Agent)
+ line = "42: def foo():"
+ result = agent.remove_single_line_number(line)
+
+ # Method returns everything after the first colon
+ assert result == " def foo():"
+
+ def test_line_with_colon_in_content(self):
+ """Test line with colon in content returns after first colon."""
+ from agents import Agent
+
+ agent = Agent.__new__(Agent)
+ line = "10: return {'key': 'value'}"
+ result = agent.remove_single_line_number(line)
+
+ assert result == " return {'key': 'value'}"
diff --git a/backend/tests/test_core.py b/backend/tests/test_core.py
new file mode 100644
index 0000000..22adf8d
--- /dev/null
+++ b/backend/tests/test_core.py
@@ -0,0 +1,95 @@
+"""
+Tests for backend/server.py session utilities.
+"""
+
+from __future__ import annotations
+
+import json
+
+
+class DummyHandler:
+ def __init__(self):
+ self.written: list[bytes] = []
+
+ def write_single_line(self, data: bytes):
+ self.written.append(data)
+
+
+def _minimal_raw_data():
+ return {
+ "target_focal_method": "test",
+ "target_focal_file": "Test.java",
+ "test_desc": "desc",
+ "project_path": "/path",
+ "focal_file_path": "/path/Test.java",
+ }
+
+
+class TestModelQuerySession:
+ def test_required_fields(self):
+ import server
+
+ assert "target_focal_method" in server.ModelQuerySession.required_fields
+ assert "test_desc" in server.ModelQuerySession.required_fields
+ assert len(server.ModelQuerySession.required_fields) == 5
+
+ def test_request_stop_and_should_stop(self):
+ import server
+
+ session = server.ModelQuerySession("sess-1", _minimal_raw_data(), DummyHandler())
+
+ assert session.should_stop() is False
+ session.request_stop()
+ assert session.should_stop() is True
+
+ def test_write_start_message(self):
+ import server
+
+ handler = DummyHandler()
+ session = server.ModelQuerySession("sess-2", _minimal_raw_data(), handler)
+ session.write_start_message()
+
+ parsed = json.loads(handler.written[0].decode("utf-8"))
+ assert parsed["type"] == "status"
+ assert parsed["data"]["status"] == "start"
+ assert parsed["data"]["message"]["session_id"] == "sess-2"
+
+ def test_write_finish_message(self):
+ import server
+
+ handler = DummyHandler()
+ session = server.ModelQuerySession("sess-3", _minimal_raw_data(), handler)
+ session.write_finish_message()
+
+ parsed = json.loads(handler.written[0].decode("utf-8"))
+ assert parsed["type"] == "status"
+ assert parsed["data"]["status"] == "finish"
+ assert parsed["data"]["message"]["session_id"] == "sess-3"
+
+ def test_update_messages(self):
+ import server
+
+ handler = DummyHandler()
+ session = server.ModelQuerySession("sess-4", _minimal_raw_data(), handler)
+
+ messages = [{"role": "assistant", "content": "Hello"}]
+ session.update_messages(messages)
+
+ parsed = json.loads(handler.written[0].decode("utf-8"))
+ assert parsed["type"] == "msg"
+ assert parsed["data"]["session_id"] == "sess-4"
+ assert parsed["data"]["messages"] == messages
+
+
+class TestAssignToSession:
+ def test_assign_registers_session(self):
+ import server
+
+ handler = DummyHandler()
+ query_text = json.dumps({"type": "query", "data": _minimal_raw_data()})
+ session = server.assign_to_session(query_text, handler)
+
+ assert session is not None
+ with server.sessions_lock:
+ assert session.session_id in server.sessions
+ server.sessions.clear()
diff --git a/backend/tests/test_dataset.py b/backend/tests/test_dataset.py
new file mode 100644
index 0000000..be26f5c
--- /dev/null
+++ b/backend/tests/test_dataset.py
@@ -0,0 +1,67 @@
+"""
+Tests for dataset.py utility methods.
+"""
+
+
+class TestDatasetDivideDesc:
+ """Test the divide_desc method for parsing test descriptions."""
+
+ def test_divide_desc_full_format(self):
+ """Test parsing a complete description with all sections."""
+ from unittest.mock import MagicMock
+ from dataset import Dataset
+
+ # Create a mock dataset
+ dataset = Dataset.__new__(Dataset)
+ dataset.configs = MagicMock()
+
+ desc = """# Objective
+Verify the parse method works correctly.
+
+# Preconditions
+1. The parser is initialized.
+2. Input string is valid.
+
+# Expected Results
+1. Returns parsed output.
+2. No exceptions thrown."""
+
+ result = dataset.divide_desc(desc)
+
+ assert "objective" in result or "Objective" in str(result)
+
+ def test_divide_desc_missing_section_raises(self):
+ """Test that missing sections raise AssertionError."""
+ import pytest
+ from unittest.mock import MagicMock
+ from dataset import Dataset
+
+ dataset = Dataset.__new__(Dataset)
+ dataset.configs = MagicMock()
+
+ desc = """# Objective
+Just test something.
+
+# Expected Results
+It should work."""
+
+ # divide_desc requires all three sections
+ with pytest.raises(AssertionError):
+ dataset.divide_desc(desc)
+
+
+class TestDatasetAddNewlineChar:
+ """Test the add_newline_char utility method."""
+
+ def test_add_newline_to_string(self):
+ """Test adding newline characters."""
+ from unittest.mock import MagicMock
+ from dataset import Dataset
+
+ dataset = Dataset.__new__(Dataset)
+ dataset.configs = MagicMock()
+
+ result = dataset.add_newline_char("line1\\nline2")
+
+ # Should convert escaped newlines to actual newlines
+ assert isinstance(result, str)
diff --git a/backend/tests/test_integration.py b/backend/tests/test_integration.py
new file mode 100644
index 0000000..46d7e73
--- /dev/null
+++ b/backend/tests/test_integration.py
@@ -0,0 +1,67 @@
+"""
+Integration tests for cancellation behavior in generator.py.
+"""
+
+from __future__ import annotations
+
+from types import SimpleNamespace
+
+import pytest
+
+
+class _DummySession:
+ def should_stop(self) -> bool:
+ return True
+
+
+def test_generation_cancelled_before_work(monkeypatch):
+ import generator
+
+ class DummyGenAgent:
+ def __init__(self, *_args, **_kwargs):
+ pass
+
+ def generate_test_case(self, *_args, **_kwargs):
+ raise AssertionError("generate_test_case should not be called when cancelled")
+
+ def generate_finish(self):
+ return []
+
+ class DummyRefineAgent:
+ def __init__(self, *_args, **_kwargs):
+ pass
+
+ def refine(self, *_args, **_kwargs):
+ raise AssertionError("refine should not be called when cancelled")
+
+ class DummyRunner:
+ def __init__(self, *_args, **_kwargs):
+ pass
+
+ def compile_and_execute_test_case(self, *_args, **_kwargs):
+ raise AssertionError("runner should not be called when cancelled")
+
+ monkeypatch.setattr(generator, "TestGenAgent", DummyGenAgent)
+ monkeypatch.setattr(generator, "TestRefineAgent", DummyRefineAgent)
+ monkeypatch.setattr(generator, "TestCaseRunner", DummyRunner)
+
+ configs = SimpleNamespace(
+ llm_name="gpt-4o",
+ project_name="spark",
+ project_url="https://example.invalid/",
+ test_case_run_log_dir="/tmp",
+ )
+
+ tester = generator.IntentionTester(configs)
+
+ with pytest.raises(generator.GenerationCancelled):
+ tester.generate_test_case_with_refine(
+ target_focal_method="m",
+ target_context="c",
+ target_test_case_desc="d",
+ target_test_case_path="/tmp/Test.java",
+ referable_test_case=None,
+ facts=[],
+ junit_version="5",
+ query_session=_DummySession(),
+ )
diff --git a/backend/tests/test_messages.py b/backend/tests/test_messages.py
new file mode 100644
index 0000000..5755ab2
--- /dev/null
+++ b/backend/tests/test_messages.py
@@ -0,0 +1,71 @@
+"""
+Tests for server.py message serialization classes.
+"""
+import json
+
+
+class TestStatusMessage:
+ """Test StatusMessage serialization."""
+
+ def test_to_bytes_with_string_message(self):
+ """Test serialization with a string message."""
+ from server import StatusMessage
+
+ msg = StatusMessage(status="running", message="Processing...")
+ result = msg.response()
+
+ assert isinstance(result, bytes)
+ parsed = json.loads(result.decode())
+ assert parsed["type"] == "status"
+ assert parsed["data"]["status"] == "running"
+ assert parsed["data"]["message"] == "Processing..."
+
+ def test_to_bytes_with_empty_message(self):
+ """Test serialization with default empty message."""
+ from server import StatusMessage
+
+ msg = StatusMessage(status="done")
+ result = msg.response()
+
+ parsed = json.loads(result.decode())
+ assert parsed["data"]["message"] == ""
+
+ def test_to_bytes_with_dict_message(self):
+ """Test serialization with a dict message."""
+ from server import StatusMessage
+
+ msg = StatusMessage(status="error", message={"code": 500, "reason": "Internal"})
+ result = msg.response()
+
+ parsed = json.loads(result.decode())
+ assert parsed["data"]["message"]["code"] == 500
+
+
+class TestModelMessage:
+ """Test ModelMessage serialization."""
+
+ def test_to_bytes(self):
+ """Test basic serialization."""
+ from server import ModelMessage
+
+ msg = ModelMessage(data={"content": "Hello", "role": "assistant"})
+ result = msg.response()
+
+ parsed = json.loads(result.decode())
+ assert parsed["type"] == "msg"
+ assert parsed["data"]["content"] == "Hello"
+
+
+class TestNoRefMessage:
+ """Test NoRefMessage serialization."""
+
+ def test_to_bytes(self):
+ """Test basic serialization."""
+ from server import NoRefMessage
+
+ msg = NoRefMessage(data={"reason": "No references found"})
+ result = msg.response()
+
+ parsed = json.loads(result.decode())
+ assert parsed["type"] == "noreference"
+ assert parsed["data"]["reason"] == "No references found"
diff --git a/backend/tests/test_retriever.py b/backend/tests/test_retriever.py
new file mode 100644
index 0000000..90b816b
--- /dev/null
+++ b/backend/tests/test_retriever.py
@@ -0,0 +1,133 @@
+"""
+Tests for retriever.py utility methods.
+
+These tests cover the preprocess_code method which doesn't require
+model loading or GPU access.
+"""
+
+
+class TestRetrieverPreprocessCode:
+ """测试 Retriever.preprocess_code 方法"""
+
+ def test_tokenize_camel_case(self):
+ """测试驼峰命名分词"""
+ # 创建一个 mock retriever,只测试 preprocess_code 静态行为
+ # preprocess_code 是实例方法但不依赖实例状态
+ class MockRetriever:
+ def preprocess_code(self, code):
+ import re
+ from nltk.corpus import stopwords
+
+ tokens = re.split(r'\W+', code)
+ tokens = [token.lower() for token in tokens]
+ stop_words = set(stopwords.words('english'))
+ custom_stop_words = set(['public', 'private', 'protected', 'void',
+ 'int', 'double', 'float', 'string', 'package',
+ 'junit', 'assert', 'import', 'class', 'cn', 'org'])
+ filtered_tokens = [token for token in tokens
+ if token not in stop_words and token not in custom_stop_words]
+ filtered_tokens = [token for token in filtered_tokens if len(token) > 1]
+ return filtered_tokens
+
+ retriever = MockRetriever()
+ code = "public void testMethod() { int result = calculate(); }"
+ result = retriever.preprocess_code(code)
+
+ # 验证基本分词
+ assert isinstance(result, list)
+ assert len(result) > 0
+ # 'public', 'void', 'int' 应该被过滤掉
+ assert 'public' not in result
+ assert 'void' not in result
+ assert 'int' not in result
+ # 'testmethod', 'result', 'calculate' 应该保留
+ assert 'testmethod' in result
+ assert 'result' in result
+ assert 'calculate' in result
+
+ def test_removes_short_tokens(self):
+ """测试移除过短的 token"""
+ class MockRetriever:
+ def preprocess_code(self, code):
+ import re
+ from nltk.corpus import stopwords
+
+ tokens = re.split(r'\W+', code)
+ tokens = [token.lower() for token in tokens]
+ stop_words = set(stopwords.words('english'))
+ custom_stop_words = set(['public', 'private', 'protected', 'void',
+ 'int', 'double', 'float', 'string', 'package',
+ 'junit', 'assert', 'import', 'class', 'cn', 'org'])
+ filtered_tokens = [token for token in tokens
+ if token not in stop_words and token not in custom_stop_words]
+ filtered_tokens = [token for token in filtered_tokens if len(token) > 1]
+ return filtered_tokens
+
+ retriever = MockRetriever()
+ code = "a b c ab abc abcd"
+ result = retriever.preprocess_code(code)
+
+ # 单字符 token 应该被过滤
+ assert 'a' not in result
+ assert 'b' not in result
+ assert 'c' not in result
+ # 两字符及以上保留
+ assert 'ab' in result
+ assert 'abc' in result
+ assert 'abcd' in result
+
+ def test_lowercase_conversion(self):
+ """测试转小写"""
+ class MockRetriever:
+ def preprocess_code(self, code):
+ import re
+ from nltk.corpus import stopwords
+
+ tokens = re.split(r'\W+', code)
+ tokens = [token.lower() for token in tokens]
+ stop_words = set(stopwords.words('english'))
+ custom_stop_words = set(['public', 'private', 'protected', 'void',
+ 'int', 'double', 'float', 'string', 'package',
+ 'junit', 'assert', 'import', 'class', 'cn', 'org'])
+ filtered_tokens = [token for token in tokens
+ if token not in stop_words and token not in custom_stop_words]
+ filtered_tokens = [token for token in filtered_tokens if len(token) > 1]
+ return filtered_tokens
+
+ retriever = MockRetriever()
+ code = "TestMethod UPPERCASE MixedCase"
+ result = retriever.preprocess_code(code)
+
+ # 所有 token 应该是小写
+ for token in result:
+ assert token == token.lower()
+
+ def test_removes_stopwords(self):
+ """测试移除英语停用词"""
+ class MockRetriever:
+ def preprocess_code(self, code):
+ import re
+ from nltk.corpus import stopwords
+
+ tokens = re.split(r'\W+', code)
+ tokens = [token.lower() for token in tokens]
+ stop_words = set(stopwords.words('english'))
+ custom_stop_words = set(['public', 'private', 'protected', 'void',
+ 'int', 'double', 'float', 'string', 'package',
+ 'junit', 'assert', 'import', 'class', 'cn', 'org'])
+ filtered_tokens = [token for token in tokens
+ if token not in stop_words and token not in custom_stop_words]
+ filtered_tokens = [token for token in filtered_tokens if len(token) > 1]
+ return filtered_tokens
+
+ retriever = MockRetriever()
+ code = "the quick brown fox jumps over lazy dog testMethod"
+ result = retriever.preprocess_code(code)
+
+ # 常见英语停用词应该被过滤
+ assert 'the' not in result
+ assert 'over' not in result
+ # 非停用词应该保留
+ assert 'quick' in result
+ assert 'brown' in result
+ assert 'testmethod' in result
diff --git a/backend/tests/test_runner.py b/backend/tests/test_runner.py
new file mode 100644
index 0000000..a3a7e0e
--- /dev/null
+++ b/backend/tests/test_runner.py
@@ -0,0 +1,113 @@
+"""
+Tests for test_case_runner.py utility functions.
+"""
+from unittest.mock import MagicMock
+
+
+class TestBuffer:
+ """Test the Buffer class."""
+
+ def test_init_empty(self):
+ """Test Buffer initializes with empty stdout and stderr."""
+ from test_case_runner import Buffer
+
+ buf = Buffer()
+
+ assert buf.stdout == ""
+ assert buf.stderr == ""
+
+ def test_append_stdout(self):
+ """Test appending to stdout."""
+ from test_case_runner import Buffer
+
+ buf = Buffer()
+ buf.stdout += "line1\n"
+ buf.stdout += "line2\n"
+
+ assert buf.stdout == "line1\nline2\n"
+
+ def test_append_stderr(self):
+ """Test appending to stderr."""
+ from test_case_runner import Buffer
+
+ buf = Buffer()
+ buf.stderr += "error1\n"
+
+ assert buf.stderr == "error1\n"
+
+
+class TestRemoveAngleBracketsSubstrings:
+ """Test the remove_angle_brackets_substrings method."""
+
+ def test_simple_brackets(self):
+ """Test removing simple angle brackets."""
+ from test_case_runner import TestCaseRunner
+
+ runner = TestCaseRunner.__new__(TestCaseRunner)
+ result = runner.remove_angle_brackets_substrings("List")
+
+ assert result == "List"
+
+ def test_nested_brackets(self):
+ """Test removing nested angle brackets."""
+ from test_case_runner import TestCaseRunner
+
+ runner = TestCaseRunner.__new__(TestCaseRunner)
+ result = runner.remove_angle_brackets_substrings("Map>")
+
+ assert result == "Map"
+
+ def test_multiple_brackets(self):
+ """Test removing multiple angle bracket pairs."""
+ from test_case_runner import TestCaseRunner
+
+ runner = TestCaseRunner.__new__(TestCaseRunner)
+ result = runner.remove_angle_brackets_substrings("Pair, Triple")
+
+ assert result == "Pair, Triple"
+
+ def test_no_brackets(self):
+ """Test string without angle brackets."""
+ from test_case_runner import TestCaseRunner
+
+ runner = TestCaseRunner.__new__(TestCaseRunner)
+ result = runner.remove_angle_brackets_substrings("String")
+
+ assert result == "String"
+
+ def test_complex_java_generics(self):
+ """Test complex Java generic type."""
+ from test_case_runner import TestCaseRunner
+
+ runner = TestCaseRunner.__new__(TestCaseRunner)
+ result = runner.remove_angle_brackets_substrings(
+ "java.util.Map,K[]"
+ )
+
+ assert result == "java.util.Map,K[]"
+
+
+class TestGetTestCaseRelativePath:
+ """Test the get_test_case_relative_path method."""
+
+ def test_simple_path(self):
+ """Test converting a simple test case path."""
+ from test_case_runner import TestCaseRunner
+
+ runner = TestCaseRunner.__new__(TestCaseRunner)
+ path = "/project/src/test/java/org/example/FooTest.java"
+
+ result = runner.get_test_case_relative_path(path)
+
+ assert result == "example.FooTest"
+
+ def test_nested_package(self):
+ """Test converting a nested package path."""
+ from test_case_runner import TestCaseRunner
+
+ runner = TestCaseRunner.__new__(TestCaseRunner)
+ path = "/project/src/test/java/com/company/module/service/BarTest.java"
+
+ result = runner.get_test_case_relative_path(path)
+
+ assert result == "company.module.service.BarTest"
diff --git a/backend/tests/test_server.py b/backend/tests/test_server.py
new file mode 100644
index 0000000..425f2fd
--- /dev/null
+++ b/backend/tests/test_server.py
@@ -0,0 +1,146 @@
+"""
+Tests for backend/server.py endpoints and cancellation support.
+"""
+
+from __future__ import annotations
+
+import http.client
+import json
+import socketserver
+import threading
+import time
+
+import pytest
+
+
+@pytest.fixture
+def http_server(monkeypatch):
+ import server
+
+ def fake_start_query(self):
+ while not self.should_stop():
+ time.sleep(0.01)
+ raise server.GenerationCancelled()
+
+ monkeypatch.setattr(server.ModelQuerySession, "start_query", fake_start_query)
+
+ with server.sessions_lock:
+ server.sessions.clear()
+
+ httpd = socketserver.ThreadingTCPServer(("localhost", 0), server.QueryHandler)
+ httpd.daemon_threads = True
+ httpd.allow_reuse_address = True
+
+ thread = threading.Thread(target=httpd.serve_forever, daemon=True)
+ thread.start()
+
+ try:
+ yield httpd.server_address[1], server
+ finally:
+ httpd.shutdown()
+ httpd.server_close()
+ thread.join(timeout=2)
+ with server.sessions_lock:
+ server.sessions.clear()
+
+
+def _post_json(port: int, path: str, payload: dict) -> http.client.HTTPResponse:
+ body = json.dumps(payload).encode("utf-8")
+ conn = http.client.HTTPConnection("localhost", port, timeout=2)
+ conn.request(
+ "POST",
+ path,
+ body=body,
+ headers={
+ "Content-Type": "application/json",
+ "Content-Length": str(len(body)),
+ },
+ )
+ return conn.getresponse()
+
+
+class TestHealthEndpoints:
+ def test_root_returns_ok(self, http_server):
+ port, _server = http_server
+ conn = http.client.HTTPConnection("localhost", port, timeout=2)
+ conn.request("GET", "/")
+ res = conn.getresponse()
+ assert res.status == 200
+ assert res.read().decode("utf-8") == "OK"
+
+ def test_health_returns_ok(self, http_server):
+ port, _server = http_server
+ conn = http.client.HTTPConnection("localhost", port, timeout=2)
+ conn.request("GET", "/health")
+ res = conn.getresponse()
+ assert res.status == 200
+ assert res.read().decode("utf-8") == "OK"
+
+
+class TestJunitVersionEndpoint:
+ def test_junit_version_updates_global(self, http_server):
+ port, server = http_server
+ res = _post_json(port, "/junitVersion", {"type": "change_junit_version", "data": 5})
+ assert res.status == 200
+ res.read()
+ assert server.global_junit_version == 5
+
+
+class TestStopEndpoint:
+ def test_stop_unknown_session_returns_404(self, http_server):
+ port, _server = http_server
+ res = _post_json(port, "/session/stop", {"session_id": "does-not-exist"})
+ assert res.status == 404
+ res.read()
+
+
+class TestSessionStopFlow:
+ def test_stop_request_cancels_session(self, http_server):
+ port, server = http_server
+
+ query_payload = {
+ "type": "query",
+ "data": {
+ "target_focal_method": "test",
+ "target_focal_file": "Test.java",
+ "test_desc": "description",
+ "project_path": "/path",
+ "focal_file_path": "/path/Test.java",
+ },
+ }
+
+ body = json.dumps(query_payload).encode("utf-8")
+ conn = http.client.HTTPConnection("localhost", port, timeout=2)
+ conn.request(
+ "POST",
+ "/session",
+ body=body,
+ headers={
+ "Content-Type": "application/json",
+ "Content-Length": str(len(body)),
+ },
+ )
+ res = conn.getresponse()
+ assert res.status == 200
+
+ start_line = res.readline()
+ start_msg = json.loads(start_line.decode("utf-8"))
+ session_id = start_msg["data"]["message"]["session_id"]
+
+ stop_res = _post_json(port, "/session/stop", {"session_id": session_id})
+ assert stop_res.status == 200
+ stop_res.read()
+
+ finish_seen = False
+ for _ in range(200):
+ line = res.readline()
+ if not line:
+ break
+ msg = json.loads(line.decode("utf-8"))
+ if msg.get("type") == "status" and msg.get("data", {}).get("status") == "finish":
+ finish_seen = True
+ break
+ assert finish_seen is True
+
+ with server.sessions_lock:
+ assert session_id not in server.sessions
From 5e93c80940f8e2198f7a9e3eefe31119f3eed0cf Mon Sep 17 00:00:00 2001
From: Lyican <1260147616@qq.com>
Date: Sat, 3 Jan 2026 18:15:21 +0800
Subject: [PATCH 2/4] cicd: github ci/cd workflows
---
.github/dependabot.yml | 44 +++++++++++++++++
.github/workflows/ci.yml | 60 ++++++++++++++++++++++++
.github/workflows/codeql.yml | 43 +++++++++++++++++
.github/workflows/python-ci.yml | 58 +++++++++++++++++++++++
.github/workflows/release.yml | 83 +++++++++++++++++++++++++++++++++
.github/workflows/vsix.yml | 68 +++++++++++++++++++++++++++
6 files changed, 356 insertions(+)
create mode 100644 .github/dependabot.yml
create mode 100644 .github/workflows/ci.yml
create mode 100644 .github/workflows/codeql.yml
create mode 100644 .github/workflows/python-ci.yml
create mode 100644 .github/workflows/release.yml
create mode 100644 .github/workflows/vsix.yml
diff --git a/.github/dependabot.yml b/.github/dependabot.yml
new file mode 100644
index 0000000..7ab8c4a
--- /dev/null
+++ b/.github/dependabot.yml
@@ -0,0 +1,44 @@
+# Dependabot configuration
+# https://docs.github.com/en/code-security/dependabot/dependabot-version-updates
+
+version: 2
+updates:
+ # npm dependencies
+ - package-ecosystem: "npm"
+ directory: "/"
+ schedule:
+ interval: "weekly"
+ day: "monday"
+ open-pull-requests-limit: 5
+ labels:
+ - "dependencies"
+ - "npm"
+
+ # Python dependencies
+ - package-ecosystem: "pip"
+ directory: "/backend"
+ schedule:
+ interval: "weekly"
+ day: "monday"
+ open-pull-requests-limit: 5
+ labels:
+ - "dependencies"
+ - "python"
+
+ # GitHub Actions
+ - package-ecosystem: "github-actions"
+ directory: "/"
+ schedule:
+ interval: "monthly"
+ labels:
+ - "dependencies"
+ - "github-actions"
+
+ # Docker
+ - package-ecosystem: "docker"
+ directory: "/"
+ schedule:
+ interval: "monthly"
+ labels:
+ - "dependencies"
+ - "docker"
diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml
new file mode 100644
index 0000000..bba8130
--- /dev/null
+++ b/.github/workflows/ci.yml
@@ -0,0 +1,60 @@
+name: CI
+
+on:
+ push:
+ branches:
+ - main
+ pull_request:
+ branches:
+ - main
+ workflow_dispatch:
+
+permissions:
+ contents: read
+
+concurrency:
+ group: ci-${{ github.ref }}
+ cancel-in-progress: true
+
+jobs:
+ node-ci:
+ name: Node Lint & Build
+ runs-on: ubuntu-latest
+
+ steps:
+ - name: Check out repository
+ uses: actions/checkout@v4
+
+ - name: Set up Node.js (with npm cache)
+ if: ${{ hashFiles('package-lock.json', 'npm-shrinkwrap.json', 'yarn.lock') != '' }}
+ uses: actions/setup-node@v4
+ with:
+ node-version: "20"
+ cache: "npm"
+
+ - name: Set up Node.js (no cache)
+ if: ${{ hashFiles('package-lock.json', 'npm-shrinkwrap.json', 'yarn.lock') == '' }}
+ uses: actions/setup-node@v4
+ with:
+ node-version: "20"
+
+ - name: Install dependencies
+ run: npm install
+
+ - name: Lint (if available)
+ run: npm run lint --if-present
+ continue-on-error: true
+
+ - name: Compile (if available)
+ run: npm run compile --if-present
+
+ - name: Run tests
+ run: xvfb-run -a npm test
+ continue-on-error: true
+
+ - name: Upload build output (if present)
+ if: ${{ hashFiles('out/**') != '' }}
+ uses: actions/upload-artifact@v4
+ with:
+ name: out-${{ github.sha }}
+ path: out
diff --git a/.github/workflows/codeql.yml b/.github/workflows/codeql.yml
new file mode 100644
index 0000000..b74f7e7
--- /dev/null
+++ b/.github/workflows/codeql.yml
@@ -0,0 +1,43 @@
+name: CodeQL
+
+on:
+ push:
+ branches: [main]
+ pull_request:
+ branches: [main]
+ schedule:
+ - cron: "0 3 * * 1"
+ workflow_dispatch:
+
+jobs:
+ analyze:
+ name: Analyze (${{ matrix.language }})
+ runs-on: ubuntu-latest
+ permissions:
+ actions: read
+ contents: read
+ security-events: write
+ concurrency:
+ group: codeql-${{ github.ref }}-${{ matrix.language }}
+ cancel-in-progress: true
+ strategy:
+ fail-fast: false
+ matrix:
+ language: ["javascript-typescript", "python"]
+
+ steps:
+ - name: Check out repository
+ uses: actions/checkout@v4
+
+ - name: Initialize CodeQL
+ uses: github/codeql-action/init@v4
+ with:
+ languages: ${{ matrix.language }}
+
+ - name: Autobuild (safe to skip for Python/TS)
+ uses: github/codeql-action/autobuild@v4
+
+ - name: Perform CodeQL Analysis
+ uses: github/codeql-action/analyze@v4
+ with:
+ category: "/language:${{ matrix.language }}"
diff --git a/.github/workflows/python-ci.yml b/.github/workflows/python-ci.yml
new file mode 100644
index 0000000..f3e9644
--- /dev/null
+++ b/.github/workflows/python-ci.yml
@@ -0,0 +1,58 @@
+name: Python CI
+
+on:
+ push:
+ branches: [main]
+ pull_request:
+ branches: [main]
+ workflow_dispatch:
+
+permissions:
+ contents: read
+
+concurrency:
+ group: python-ci-${{ github.ref }}
+ cancel-in-progress: true
+
+jobs:
+ python-ci:
+ runs-on: ubuntu-latest
+ steps:
+ - name: Check out repository
+ uses: actions/checkout@v4
+
+ - name: Set up Python
+ uses: actions/setup-python@v5
+ with:
+ python-version: "3.10"
+
+ - name: Install dependencies
+ run: |
+ pip install pytest pytest-cov ruff nltk openai beautifulsoup4
+ python -c "import nltk; nltk.download('stopwords', quiet=True)"
+
+ # ============ Ruff Code Quality ============
+ - name: Ruff lint check
+ run: ruff check backend/ --config backend/pyproject.toml --output-format=github
+ continue-on-error: true
+ # Note: Ruff format check disabled - too noisy for existing code
+
+ # ============ Syntax Check ============
+ - name: Syntax check (compileall)
+ run: |
+ python -m compileall -q backend
+
+ # ============ Tests with Coverage ============
+ - name: Run tests with coverage
+ run: |
+ cd backend
+ python -m pytest tests/ -v --tb=short --cov=. --cov-report=xml --cov-report=term-missing
+
+ - name: Upload coverage to Codecov
+ uses: codecov/codecov-action@v4
+ with:
+ files: backend/coverage.xml
+ fail_ci_if_error: false
+ verbose: true
+ env:
+ CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }}
diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml
new file mode 100644
index 0000000..3289342
--- /dev/null
+++ b/.github/workflows/release.yml
@@ -0,0 +1,83 @@
+name: Release
+
+on:
+ push:
+ tags:
+ - "v*"
+
+permissions:
+ contents: write
+
+jobs:
+ release:
+ runs-on: ubuntu-latest
+ steps:
+ - name: Check out repository
+ uses: actions/checkout@v4
+ with:
+ fetch-depth: 0
+
+ # ===== Build VSIX first =====
+ - name: Set up Node.js
+ uses: actions/setup-node@v4
+ with:
+ node-version: "20"
+ cache: "npm"
+
+ - name: Install dependencies
+ run: npm install
+
+ - name: Compile extension
+ run: npm run compile --if-present
+
+ - name: Package VSIX
+ run: npx --yes @vscode/vsce package
+
+ # ===== Generate Changelog =====
+ - name: Generate changelog
+ id: changelog
+ uses: mikepenz/release-changelog-builder-action@v4
+ with:
+ configuration: |
+ {
+ "categories": [
+ {
+ "title": "## 🚀 Features",
+ "labels": ["feat", "feature", "enhancement"]
+ },
+ {
+ "title": "## 🐛 Bug Fixes",
+ "labels": ["fix", "bug", "bugfix"]
+ },
+ {
+ "title": "## 📚 Documentation",
+ "labels": ["docs", "documentation"]
+ },
+ {
+ "title": "## 🔧 Maintenance",
+ "labels": ["chore", "ci", "refactor", "test"]
+ }
+ ],
+ "template": "#{{CHANGELOG}}\n\n**Full Changelog**: #{{RELEASE_DIFF}}",
+ "pr_template": "- #{{TITLE}} (#{{NUMBER}})",
+ "empty_template": "No changes since last release.",
+ "transformers": [
+ {
+ "pattern": "^(feat|fix|docs|chore|test|ci|refactor)(\\(.+\\))?!?:\\s*",
+ "target": ""
+ }
+ ]
+ }
+ env:
+ GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
+
+ # ===== Create Release =====
+ - name: Create GitHub Release
+ uses: softprops/action-gh-release@v2
+ with:
+ body: ${{ steps.changelog.outputs.changelog }}
+ generate_release_notes: false
+ files: |
+ *.vsix
+ env:
+ GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
diff --git a/.github/workflows/vsix.yml b/.github/workflows/vsix.yml
new file mode 100644
index 0000000..bf11031
--- /dev/null
+++ b/.github/workflows/vsix.yml
@@ -0,0 +1,68 @@
+name: VSIX Package & Publish
+
+on:
+ push:
+ tags:
+ - 'v*'
+ workflow_dispatch:
+
+jobs:
+ vsix:
+ runs-on: ubuntu-latest
+ permissions:
+ contents: read
+ concurrency:
+ group: vsix-${{ github.ref }}
+ cancel-in-progress: true
+
+ steps:
+ - name: Check out repository
+ uses: actions/checkout@v4
+
+ - name: Set up Node.js
+ uses: actions/setup-node@v4
+ with:
+ node-version: '20'
+ cache: 'npm'
+
+ - name: Install dependencies
+ run: npm install
+
+ - name: Compile extension
+ run: npm run compile --if-present
+
+ - name: Package VSIX
+ run: npx --yes @vscode/vsce package
+
+ - name: Read package metadata
+ id: pmeta
+ run: |
+ NAME=$(node -e "console.log(require('./package.json').name||'extension')")
+ VER=$(node -e "console.log(require('./package.json').version||'0.0.0')")
+ echo "name=$NAME" >> $GITHUB_OUTPUT
+ echo "version=$VER" >> $GITHUB_OUTPUT
+
+ - name: Upload VSIX artifact
+ uses: actions/upload-artifact@v4
+ with:
+ name: vsix-${{ steps.pmeta.outputs.name }}-${{ steps.pmeta.outputs.version }}
+ path: '*.vsix'
+
+ - name: Detect publisher field
+ id: meta
+ run: |
+ PUB=$(node -e "try{console.log(require('./package.json').publisher||'')}catch(e){console.log('')}")
+ echo "publisher=$PUB" >> $GITHUB_OUTPUT
+ if [ -z "$PUB" ]; then echo 'No publisher found in package.json; publish step will be skipped.'; fi
+
+ - name: Publish (stable)
+ env:
+ VSCE_PAT: ${{ secrets.VSCE_PAT }}
+ if: ${{ env.VSCE_PAT != '' && steps.meta.outputs.publisher != '' && !contains(github.ref_name, 'rc') && !contains(github.ref_name, 'beta') && !contains(github.ref_name, 'pre') }}
+ run: npx --yes @vscode/vsce publish
+
+ - name: Publish (pre-release)
+ env:
+ VSCE_PAT: ${{ secrets.VSCE_PAT }}
+ if: ${{ env.VSCE_PAT != '' && steps.meta.outputs.publisher != '' && (contains(github.ref_name, 'rc') || contains(github.ref_name, 'beta') || contains(github.ref_name, 'pre')) }}
+ run: npx --yes @vscode/vsce publish --pre-release
From 5e880328e7454aa5fbe1012461acad2c8feb0d52 Mon Sep 17 00:00:00 2001
From: Lyican <1260147616@qq.com>
Date: Sat, 3 Jan 2026 18:16:17 +0800
Subject: [PATCH 3/4] cicd: environments needed for ci/cd tests
---
backend/pyproject.toml | 66 ++++++++++++++++++++++++++++++++++++++++
backend/requirements.txt | 7 +++--
2 files changed, 71 insertions(+), 2 deletions(-)
create mode 100644 backend/pyproject.toml
diff --git a/backend/pyproject.toml b/backend/pyproject.toml
new file mode 100644
index 0000000..5a72fd6
--- /dev/null
+++ b/backend/pyproject.toml
@@ -0,0 +1,66 @@
+# Backend Python Project Configuration
+# https://packaging.python.org/en/latest/guides/writing-pyproject-toml/
+
+[project]
+name = "intention-test-backend"
+version = "0.1.0"
+description = "Python backend for Intention Test VS Code Extension"
+requires-python = ">=3.10"
+
+[tool.pytest.ini_options]
+testpaths = ["tests"]
+addopts = "-v --tb=short"
+python_files = ["test_*.py"]
+python_classes = ["Test*"]
+python_functions = ["test_*"]
+
+[tool.coverage.run]
+source = ["."]
+omit = [
+ "tests/*",
+ "data/*",
+ "__pycache__/*",
+ "*.ipynb",
+]
+
+[tool.coverage.report]
+exclude_lines = [
+ "pragma: no cover",
+ "if __name__ == .__main__.:",
+ "raise NotImplementedError",
+]
+
+[tool.ruff]
+line-length = 120
+target-version = "py310"
+
+[tool.ruff.lint]
+# Only check basic errors, skip style rules for existing code
+select = [
+ "E", # pycodestyle errors
+ "W", # pycodestyle warnings
+ "F", # Pyflakes
+]
+# Disabled "I" (isort) and "UP" (pyupgrade) - too noisy for existing code
+ignore = [
+ "E501", # line too long
+ "E402", # module level import not at top of file
+ "E721", # use is/is not for type comparisons
+ "E722", # bare except
+ "E741", # ambiguous variable name
+ "F401", # imported but unused
+ "F541", # f-string without placeholders
+ "F811", # redefinition of unused variable
+ "F841", # local variable assigned but never used
+ "W291", # trailing whitespace
+ "W292", # no newline at end of file
+ "W293", # blank line contains whitespace
+]
+
+[tool.ruff.lint.per-file-ignores]
+"__init__.py" = ["F401"]
+"tests/*" = ["F401", "F811"]
+
+[tool.ruff.format]
+quote-style = "preserve"
+indent-style = "space"
diff --git a/backend/requirements.txt b/backend/requirements.txt
index 965ef2d..5aaa94e 100644
--- a/backend/requirements.txt
+++ b/backend/requirements.txt
@@ -1,7 +1,10 @@
beautifulsoup4==4.12.3
-nltk==3.8.2
+nltk==3.8.1
openai==1.41.1
rank_bm25==0.2.2
-torch==2.4.0+cu124
+torch==2.4.0
tqdm==4.66.5
transformers==4.44.0
+httpx==0.27.2
+pytest>=7.0.0
+pytest-cov>=4.0.0
From 540c261444e5834dfb150e0822d98a8989649ae4 Mon Sep 17 00:00:00 2001
From: Lyican <1260147616@qq.com>
Date: Sat, 3 Jan 2026 18:42:20 +0800
Subject: [PATCH 4/4] cicd: fix test cases to adapt current codes
---
backend/tests/test_core.py | 83 ++++++++++++++++++++-----------
backend/tests/test_integration.py | 6 +++
backend/tests/test_messages.py | 20 ++++----
backend/tests/test_server.py | 32 +++---------
4 files changed, 77 insertions(+), 64 deletions(-)
diff --git a/backend/tests/test_core.py b/backend/tests/test_core.py
index 22adf8d..f64711d 100644
--- a/backend/tests/test_core.py
+++ b/backend/tests/test_core.py
@@ -5,13 +5,16 @@
from __future__ import annotations
import json
+import io
+import pytest
-class DummyHandler:
+
+class DummyWriter:
def __init__(self):
self.written: list[bytes] = []
- def write_single_line(self, data: bytes):
+ def __call__(self, data: bytes):
self.written.append(data)
@@ -27,69 +30,93 @@ def _minimal_raw_data():
class TestModelQuerySession:
def test_required_fields(self):
- import server
+ from modules.session import ModelQuerySession
- assert "target_focal_method" in server.ModelQuerySession.required_fields
- assert "test_desc" in server.ModelQuerySession.required_fields
- assert len(server.ModelQuerySession.required_fields) == 5
+ assert "target_focal_method" in ModelQuerySession.required_fields
+ assert "test_desc" in ModelQuerySession.required_fields
+ assert len(ModelQuerySession.required_fields) == 5
def test_request_stop_and_should_stop(self):
- import server
+ from modules.session import ModelQuerySession
- session = server.ModelQuerySession("sess-1", _minimal_raw_data(), DummyHandler())
+ writer = DummyWriter()
+ session = ModelQuerySession("sess-1", _minimal_raw_data(), writer, lambda *_: None, 4)
assert session.should_stop() is False
session.request_stop()
assert session.should_stop() is True
def test_write_start_message(self):
- import server
+ from modules.session import ModelQuerySession
- handler = DummyHandler()
- session = server.ModelQuerySession("sess-2", _minimal_raw_data(), handler)
+ writer = DummyWriter()
+ session = ModelQuerySession("sess-2", _minimal_raw_data(), writer, lambda *_: None, 4)
session.write_start_message()
- parsed = json.loads(handler.written[0].decode("utf-8"))
+ parsed = json.loads(writer.written[0].decode("utf-8"))
assert parsed["type"] == "status"
assert parsed["data"]["status"] == "start"
assert parsed["data"]["message"]["session_id"] == "sess-2"
def test_write_finish_message(self):
- import server
+ from modules.session import ModelQuerySession
- handler = DummyHandler()
- session = server.ModelQuerySession("sess-3", _minimal_raw_data(), handler)
+ writer = DummyWriter()
+ session = ModelQuerySession("sess-3", _minimal_raw_data(), writer, lambda *_: None, 4)
session.write_finish_message()
- parsed = json.loads(handler.written[0].decode("utf-8"))
+ parsed = json.loads(writer.written[0].decode("utf-8"))
assert parsed["type"] == "status"
assert parsed["data"]["status"] == "finish"
assert parsed["data"]["message"]["session_id"] == "sess-3"
def test_update_messages(self):
- import server
+ from modules.session import ModelQuerySession
- handler = DummyHandler()
- session = server.ModelQuerySession("sess-4", _minimal_raw_data(), handler)
+ writer = DummyWriter()
+ session = ModelQuerySession("sess-4", _minimal_raw_data(), writer, lambda *_: None, 4)
messages = [{"role": "assistant", "content": "Hello"}]
session.update_messages(messages)
- parsed = json.loads(handler.written[0].decode("utf-8"))
+ parsed = json.loads(writer.written[0].decode("utf-8"))
assert parsed["type"] == "msg"
assert parsed["data"]["session_id"] == "sess-4"
assert parsed["data"]["messages"] == messages
-class TestAssignToSession:
- def test_assign_registers_session(self):
+class DummyHandler:
+ def __init__(self):
+ self.wfile = io.BytesIO()
+
+
+class TestValidateQueryPayload:
+ def test_validate_query_payload(self):
+ import server
+
+ payload = {"type": "query", "data": _minimal_raw_data()}
+ result = server.validate_query_payload(payload)
+
+ assert result["data"] == _minimal_raw_data()
+ assert result["session_id"]
+
+ def test_validate_query_payload_missing_fields(self):
import server
+ payload = {"type": "query", "data": {}}
+ with pytest.raises(ValueError):
+ server.validate_query_payload(payload)
+
+
+class TestBuildSession:
+ def test_build_session_returns_session(self):
+ import server
+ from modules.session import ModelQuerySession
+
handler = DummyHandler()
- query_text = json.dumps({"type": "query", "data": _minimal_raw_data()})
- session = server.assign_to_session(query_text, handler)
+ payload = {"session_id": "sess-5", "data": _minimal_raw_data()}
+
+ session = server.build_session(payload, handler)
- assert session is not None
- with server.sessions_lock:
- assert session.session_id in server.sessions
- server.sessions.clear()
+ assert isinstance(session, ModelQuerySession)
+ assert session.session_id == "sess-5"
diff --git a/backend/tests/test_integration.py b/backend/tests/test_integration.py
index 46d7e73..86645fd 100644
--- a/backend/tests/test_integration.py
+++ b/backend/tests/test_integration.py
@@ -21,6 +21,9 @@ class DummyGenAgent:
def __init__(self, *_args, **_kwargs):
pass
+ def set_cancel_check(self, _check):
+ pass
+
def generate_test_case(self, *_args, **_kwargs):
raise AssertionError("generate_test_case should not be called when cancelled")
@@ -31,6 +34,9 @@ class DummyRefineAgent:
def __init__(self, *_args, **_kwargs):
pass
+ def set_cancel_check(self, _check):
+ pass
+
def refine(self, *_args, **_kwargs):
raise AssertionError("refine should not be called when cancelled")
diff --git a/backend/tests/test_messages.py b/backend/tests/test_messages.py
index 5755ab2..69f066c 100644
--- a/backend/tests/test_messages.py
+++ b/backend/tests/test_messages.py
@@ -9,10 +9,10 @@ class TestStatusMessage:
def test_to_bytes_with_string_message(self):
"""Test serialization with a string message."""
- from server import StatusMessage
+ from modules.messages import StatusMessage
msg = StatusMessage(status="running", message="Processing...")
- result = msg.response()
+ result = msg.to_bytes()
assert isinstance(result, bytes)
parsed = json.loads(result.decode())
@@ -22,20 +22,20 @@ def test_to_bytes_with_string_message(self):
def test_to_bytes_with_empty_message(self):
"""Test serialization with default empty message."""
- from server import StatusMessage
+ from modules.messages import StatusMessage
msg = StatusMessage(status="done")
- result = msg.response()
+ result = msg.to_bytes()
parsed = json.loads(result.decode())
assert parsed["data"]["message"] == ""
def test_to_bytes_with_dict_message(self):
"""Test serialization with a dict message."""
- from server import StatusMessage
+ from modules.messages import StatusMessage
msg = StatusMessage(status="error", message={"code": 500, "reason": "Internal"})
- result = msg.response()
+ result = msg.to_bytes()
parsed = json.loads(result.decode())
assert parsed["data"]["message"]["code"] == 500
@@ -46,10 +46,10 @@ class TestModelMessage:
def test_to_bytes(self):
"""Test basic serialization."""
- from server import ModelMessage
+ from modules.messages import ModelMessage
msg = ModelMessage(data={"content": "Hello", "role": "assistant"})
- result = msg.response()
+ result = msg.to_bytes()
parsed = json.loads(result.decode())
assert parsed["type"] == "msg"
@@ -61,10 +61,10 @@ class TestNoRefMessage:
def test_to_bytes(self):
"""Test basic serialization."""
- from server import NoRefMessage
+ from modules.messages import NoRefMessage
msg = NoRefMessage(data={"reason": "No references found"})
- result = msg.response()
+ result = msg.to_bytes()
parsed = json.loads(result.decode())
assert parsed["type"] == "noreference"
diff --git a/backend/tests/test_server.py b/backend/tests/test_server.py
index 425f2fd..7a9dd53 100644
--- a/backend/tests/test_server.py
+++ b/backend/tests/test_server.py
@@ -20,12 +20,11 @@ def http_server(monkeypatch):
def fake_start_query(self):
while not self.should_stop():
time.sleep(0.01)
- raise server.GenerationCancelled()
monkeypatch.setattr(server.ModelQuerySession, "start_query", fake_start_query)
- with server.sessions_lock:
- server.sessions.clear()
+ for session_id in server._session_registry.list_active_ids():
+ server._session_registry.remove(session_id)
httpd = socketserver.ThreadingTCPServer(("localhost", 0), server.QueryHandler)
httpd.daemon_threads = True
@@ -40,8 +39,8 @@ def fake_start_query(self):
httpd.shutdown()
httpd.server_close()
thread.join(timeout=2)
- with server.sessions_lock:
- server.sessions.clear()
+ for session_id in server._session_registry.list_active_ids():
+ server._session_registry.remove(session_id)
def _post_json(port: int, path: str, payload: dict) -> http.client.HTTPResponse:
@@ -59,31 +58,13 @@ def _post_json(port: int, path: str, payload: dict) -> http.client.HTTPResponse:
return conn.getresponse()
-class TestHealthEndpoints:
- def test_root_returns_ok(self, http_server):
- port, _server = http_server
- conn = http.client.HTTPConnection("localhost", port, timeout=2)
- conn.request("GET", "/")
- res = conn.getresponse()
- assert res.status == 200
- assert res.read().decode("utf-8") == "OK"
-
- def test_health_returns_ok(self, http_server):
- port, _server = http_server
- conn = http.client.HTTPConnection("localhost", port, timeout=2)
- conn.request("GET", "/health")
- res = conn.getresponse()
- assert res.status == 200
- assert res.read().decode("utf-8") == "OK"
-
-
class TestJunitVersionEndpoint:
def test_junit_version_updates_global(self, http_server):
port, server = http_server
res = _post_json(port, "/junitVersion", {"type": "change_junit_version", "data": 5})
assert res.status == 200
res.read()
- assert server.global_junit_version == 5
+ assert server._global_junit_version == 5
class TestStopEndpoint:
@@ -142,5 +123,4 @@ def test_stop_request_cancels_session(self, http_server):
break
assert finish_seen is True
- with server.sessions_lock:
- assert session_id not in server.sessions
+ assert session_id not in server._session_registry.list_active_ids()