198 lines
6.9 KiB
Python
198 lines
6.9 KiB
Python
"""Regression tests for branch-shared RAG behavior."""
|
|
|
|
import ast
|
|
import json
|
|
import os
|
|
import pathlib
|
|
import unittest
|
|
from datetime import datetime
|
|
from unittest.mock import patch
|
|
|
|
from core.cache import RAGCacheManager
|
|
from core.intent_analyzer import IntentAnalyzer
|
|
from exam_pkg.generator import QuestionGenerator, analyze_document_for_exam
|
|
from exam_pkg.grader import grade_fill_blank
|
|
from knowledge.sync import ChangeType, DocumentChange, KnowledgeSyncService
|
|
|
|
|
|
class CacheInvalidationTests(unittest.TestCase):
|
|
def test_kb_version_change_clears_rerank_scores(self):
|
|
cache = RAGCacheManager()
|
|
cache.set_rerank_scores("问题", ["doc-1"], [0.91])
|
|
self.assertEqual({"doc-1": 0.91}, cache.get_rerank_scores("问题", ["doc-1"]))
|
|
|
|
cache.increment_kb_version("public")
|
|
|
|
self.assertIsNone(cache.get_rerank_scores("问题", ["doc-1"]))
|
|
|
|
|
|
class IntentCacheTests(unittest.TestCase):
|
|
def test_unknown_semantic_cache_type_is_not_used_as_intent(self):
|
|
class FakeEmbeddingModel:
|
|
def encode(self, _text):
|
|
return [1.0, 0.0]
|
|
|
|
class FakeCache:
|
|
def __init__(self):
|
|
self.saved = None
|
|
|
|
def get(self, _embedding):
|
|
return {
|
|
"cache_type": "future_cache_type",
|
|
"rewritten_query": "错误的缓存结果",
|
|
}
|
|
|
|
def set(self, _embedding, value):
|
|
self.saved = value
|
|
|
|
def get_stats(self):
|
|
return {}
|
|
|
|
analyzer = IntentAnalyzer(model="mimo-v2.5")
|
|
analyzer._embedding_model = FakeEmbeddingModel()
|
|
analyzer._cache = FakeCache()
|
|
analyzer._client = object()
|
|
|
|
llm_result = json.dumps(
|
|
{
|
|
"rewritten_query": "来自 LLM 的新结果",
|
|
"use_context": False,
|
|
"need_retrieval": True,
|
|
"reason": "新问题",
|
|
"sub_queries": ["来自 LLM 的新结果"],
|
|
"intent": "factual",
|
|
},
|
|
ensure_ascii=False,
|
|
)
|
|
with patch("core.llm_utils.call_llm", return_value=llm_result):
|
|
result = analyzer.analyze("当前问题", [])
|
|
|
|
self.assertEqual("来自 LLM 的新结果", result.rewritten_query)
|
|
self.assertEqual("intent_analysis", analyzer._cache.saved["cache_type"])
|
|
|
|
|
|
class ExamRegressionTests(unittest.TestCase):
|
|
def test_analysis_prompt_respects_requested_max_total(self):
|
|
captured = {}
|
|
|
|
def fake_call(_self, prompt):
|
|
captured["prompt"] = prompt
|
|
return json.dumps(
|
|
{
|
|
"total_knowledge_points": 2,
|
|
"suitable_types": ["single_choice"],
|
|
"question_types": {
|
|
"single_choice": 3,
|
|
"multiple_choice": 0,
|
|
"true_false": 0,
|
|
"fill_blank": 0,
|
|
"subjective": 0,
|
|
},
|
|
"reason": "测试",
|
|
},
|
|
ensure_ascii=False,
|
|
)
|
|
|
|
chunks = [{"content": "知识点" * 80, "section": "第一章"}]
|
|
with patch.object(QuestionGenerator, "_call_llm", fake_call):
|
|
result = analyze_document_for_exam(chunks, max_total=3)
|
|
|
|
self.assertIn("所有数量之和不要超过 3", captured["prompt"])
|
|
self.assertLessEqual(sum(result["question_types"].values()), 3)
|
|
|
|
def test_fill_blank_rejects_non_list_student_answer(self):
|
|
result = grade_fill_blank(
|
|
{
|
|
"question_id": "q1",
|
|
"content": {"answer": [["正确答案"]], "data": {"blank_count": 1}},
|
|
"student_answer": "正确答案",
|
|
"max_score": 2,
|
|
}
|
|
)
|
|
self.assertEqual("failed", result["grading_status"])
|
|
self.assertEqual(0, result["score"])
|
|
|
|
def test_fill_blank_rejects_non_string_items(self):
|
|
result = grade_fill_blank(
|
|
{
|
|
"question_id": "q1",
|
|
"content": {"answer": [["正确答案"]], "data": {"blank_count": 1}},
|
|
"student_answer": [{"text": "正确答案"}],
|
|
"max_score": 2,
|
|
}
|
|
)
|
|
self.assertEqual("failed", result["grading_status"])
|
|
self.assertIn("第 1 项", result["details"]["error"])
|
|
|
|
|
|
class SyncInvalidationTests(unittest.TestCase):
|
|
def test_document_delete_clears_semantic_cache(self):
|
|
class FakeDb:
|
|
def delete_document_hash(self, _document_id):
|
|
return None
|
|
|
|
class FakeKbManager:
|
|
def delete_document(self, _kb_name, _document_name):
|
|
return 1
|
|
|
|
class FakeSemanticCache:
|
|
def __init__(self):
|
|
self.cleared = False
|
|
|
|
def clear(self):
|
|
self.cleared = True
|
|
|
|
service = object.__new__(KnowledgeSyncService)
|
|
service.documents_path = "unused"
|
|
service.db = FakeDb()
|
|
semantic_cache = FakeSemanticCache()
|
|
change = DocumentChange(
|
|
document_id="public/example.pdf",
|
|
document_name="example.pdf",
|
|
change_type=ChangeType.DELETED,
|
|
old_hash="old",
|
|
new_hash=None,
|
|
change_time=datetime.now(),
|
|
)
|
|
|
|
with (
|
|
patch("knowledge.manager.get_kb_manager", return_value=FakeKbManager()),
|
|
patch("knowledge.sync.CACHE_AVAILABLE", False),
|
|
patch("core.semantic_cache.get_semantic_cache", return_value=semantic_cache),
|
|
):
|
|
self.assertTrue(service.process_change(change))
|
|
|
|
self.assertTrue(semantic_cache.cleared)
|
|
|
|
|
|
class ConfigTemplateTests(unittest.TestCase):
|
|
def test_config_example_defines_all_direct_imports(self):
|
|
root = pathlib.Path(__file__).resolve().parents[1]
|
|
template_tree = ast.parse((root / "config.example.py").read_text(encoding="utf-8"))
|
|
defined = {
|
|
node.name
|
|
for node in template_tree.body
|
|
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef))
|
|
}
|
|
for node in template_tree.body:
|
|
if isinstance(node, ast.Assign):
|
|
defined.update(target.id for target in node.targets if isinstance(target, ast.Name))
|
|
elif isinstance(node, ast.AnnAssign) and isinstance(node.target, ast.Name):
|
|
defined.add(node.target.id)
|
|
|
|
imported = set()
|
|
excluded = {"venv", ".git", ".data", ".recovery", "__pycache__"}
|
|
for path in root.rglob("*.py"):
|
|
if any(part in excluded for part in path.parts):
|
|
continue
|
|
tree = ast.parse(path.read_text(encoding="utf-8-sig"))
|
|
for node in ast.walk(tree):
|
|
if isinstance(node, ast.ImportFrom) and node.module == "config":
|
|
imported.update(alias.name for alias in node.names)
|
|
|
|
self.assertEqual([], sorted(imported - defined))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|