diff --git a/tests/unit/test_oneclick_gen_p0_fixes.py b/tests/unit/test_oneclick_gen_p0_fixes.py index 6914fd7e6..4bc20f178 100755 --- a/tests/unit/test_oneclick_gen_p0_fixes.py +++ b/tests/unit/test_oneclick_gen_p0_fixes.py @@ -241,11 +241,27 @@ class TestP1Validations: mock_template.is_active = True session = MagicMock() - mock_session = MagicMock() - session.query.return_value = mock_session - filter_result = MagicMock() - mock_session.filter.return_value = filter_result - filter_result.first.return_value = mock_template + + # EditTemplateModel 查询返回 None(走旧模板系统 fallback) + edit_query = MagicMock() + edit_filter = MagicMock() + edit_query.filter.return_value = edit_filter + edit_filter.first.return_value = None + + # TemplateModel 查询返回 mock_template + old_query = MagicMock() + old_filter = MagicMock() + old_query.filter.return_value = old_filter + old_filter.first.return_value = mock_template + + def _query_side_effect(model): + # 根据 model 类型返回不同的 query mock + name = getattr(model, "__name__", "") + if "EditTemplate" in name: + return edit_query + return old_query + + session.query.side_effect = _query_side_effect with patch("worker_app.tasks.generation.SessionLocal", return_value=session): _validate_template_exists("tmpl_001") # 不抛异常 @@ -255,11 +271,16 @@ class TestP1Validations: from worker_app.tasks.generation import _validate_template_exists session = MagicMock() - mock_session = MagicMock() - session.query.return_value = mock_session - filter_result = MagicMock() - mock_session.filter.return_value = filter_result - filter_result.first.return_value = None + + # 两个系统查询都返回 None + def _query_side_effect(model): + q = MagicMock() + f = MagicMock() + q.filter.return_value = f + f.first.return_value = None + return q + + session.query.side_effect = _query_side_effect with patch("worker_app.tasks.generation.SessionLocal", return_value=session): with pytest.raises(ValueError, match="模板不存在"): @@ -431,11 +452,26 @@ class TestTemplatePlanConfigLoading: def _mock_session(self, template): session = MagicMock() - mock_query = MagicMock() - session.query.return_value = mock_query - filter_result = MagicMock() - mock_query.filter.return_value = filter_result - filter_result.first.return_value = template + + # EditTemplateModel 查询返回 None(走旧模板系统 fallback) + edit_query = MagicMock() + edit_filter = MagicMock() + edit_query.filter.return_value = edit_filter + edit_filter.first.return_value = None + + # TemplateModel 查询返回 template(旧模板系统) + old_query = MagicMock() + old_filter = MagicMock() + old_query.filter.return_value = old_filter + old_filter.first.return_value = template + + def _query_side_effect(model): + name = getattr(model, "__name__", "") + if "EditTemplate" in name: + return edit_query + return old_query + + session.query.side_effect = _query_side_effect return session def test_load_template_config_assembles_three_fields(self): @@ -503,11 +539,15 @@ class TestTemplatePlanConfigLoading: from worker_app.tasks.generation import _load_template_plan_config session = MagicMock() - mock_query = MagicMock() - session.query.return_value = mock_query - filter_result = MagicMock() - mock_query.filter.return_value = filter_result - filter_result.first.return_value = None + + def _query_side_effect(model): + q = MagicMock() + f = MagicMock() + q.filter.return_value = f + f.first.return_value = None + return q + + session.query.side_effect = _query_side_effect with patch("worker_app.tasks.generation.SessionLocal", return_value=session): result = _load_template_plan_config("tmpl_nonexist")