Underground-Digital/Workflow-Engine
0
1import os2import time3import uuid4from typing import Optional5from unittest.mock import MagicMock6 7import pytest8 9from core.app.entities.app_invoke_entities import InvokeFrom, ModelConfigWithCredentialsEntity10from core.entities.provider_configuration import ProviderConfiguration, ProviderModelBundle11from core.entities.provider_entities import CustomConfiguration, CustomProviderConfiguration, SystemConfiguration12from core.model_manager import ModelInstance13from core.model_runtime.entities.model_entities import ModelType14from core.model_runtime.model_providers.model_provider_factory import ModelProviderFactory15from core.workflow.entities.variable_pool import VariablePool16from core.workflow.enums import SystemVariableKey17from core.workflow.graph_engine.entities.graph import Graph18from core.workflow.graph_engine.entities.graph_init_params import GraphInitParams19from core.workflow.graph_engine.entities.graph_runtime_state import GraphRuntimeState20from core.workflow.nodes.parameter_extractor.parameter_extractor_node import ParameterExtractorNode21from extensions.ext_database import db22from models.enums import UserFrom23from models.provider import ProviderType24 25"""FOR MOCK FIXTURES, DO NOT REMOVE"""26from models.workflow import WorkflowNodeExecutionStatus, WorkflowType27from tests.integration_tests.model_runtime.__mock.anthropic import setup_anthropic_mock28from tests.integration_tests.model_runtime.__mock.openai import setup_openai_mock29 30 31def get_mocked_fetch_model_config(32 provider: str,33 model: str,34 mode: str,35 credentials: dict,36):37 provider_instance = ModelProviderFactory().get_provider_instance(provider)38 model_type_instance = provider_instance.get_model_instance(ModelType.LLM)39 provider_model_bundle = ProviderModelBundle(40 configuration=ProviderConfiguration(41 tenant_id="1",42 provider=provider_instance.get_provider_schema(),43 preferred_provider_type=ProviderType.CUSTOM,44 using_provider_type=ProviderType.CUSTOM,45 system_configuration=SystemConfiguration(enabled=False),46 custom_configuration=CustomConfiguration(provider=CustomProviderConfiguration(credentials=credentials)),47 model_settings=[],48 ),49 provider_instance=provider_instance,50 model_type_instance=model_type_instance,51 )52 model_instance = ModelInstance(provider_model_bundle=provider_model_bundle, model=model)53 model_schema = model_type_instance.get_model_schema(model)54 assert model_schema is not None55 model_config = ModelConfigWithCredentialsEntity(56 model=model,57 provider=provider,58 mode=mode,59 credentials=credentials,60 parameters={},61 model_schema=model_schema,62 provider_model_bundle=provider_model_bundle,63 )64 65 return MagicMock(return_value=(model_instance, model_config))66 67 68def get_mocked_fetch_memory(memory_text: str):69 class MemoryMock:70 def get_history_prompt_text(71 self,72 human_prefix: str = "Human",73 ai_prefix: str = "Assistant",74 max_token_limit: int = 2000,75 message_limit: Optional[int] = None,76 ):77 return memory_text78 79 return MagicMock(return_value=MemoryMock())80 81 82def init_parameter_extractor_node(config: dict):83 graph_config = {84 "edges": [85 {86 "id": "start-source-next-target",87 "source": "start",88 "target": "llm",89 },90 ],91 "nodes": [{"data": {"type": "start"}, "id": "start"}, config],92 }93 94 graph = Graph.init(graph_config=graph_config)95 96 init_params = GraphInitParams(97 tenant_id="1",98 app_id="1",99 workflow_type=WorkflowType.WORKFLOW,100 workflow_id="1",101 graph_config=graph_config,102 user_id="1",103 user_from=UserFrom.ACCOUNT,104 invoke_from=InvokeFrom.DEBUGGER,105 call_depth=0,106 )107 108 # construct variable pool109 variable_pool = VariablePool(110 system_variables={111 SystemVariableKey.QUERY: "what's the weather in SF",112 SystemVariableKey.FILES: [],113 SystemVariableKey.CONVERSATION_ID: "abababa",114 SystemVariableKey.USER_ID: "aaa",115 },116 user_inputs={},117 environment_variables=[],118 conversation_variables=[],119 )120 variable_pool.add(["a", "b123", "args1"], 1)121 variable_pool.add(["a", "b123", "args2"], 2)122 123 return ParameterExtractorNode(124 id=str(uuid.uuid4()),125 graph_init_params=init_params,126 graph=graph,127 graph_runtime_state=GraphRuntimeState(variable_pool=variable_pool, start_at=time.perf_counter()),128 config=config,129 )130 131 132@pytest.mark.parametrize("setup_openai_mock", [["chat"]], indirect=True)133def test_function_calling_parameter_extractor(setup_openai_mock):134 """135 Test function calling for parameter extractor.136 """137 node = init_parameter_extractor_node(138 config={139 "id": "llm",140 "data": {141 "title": "123",142 "type": "parameter-extractor",143 "model": {"provider": "openai", "name": "gpt-3.5-turbo", "mode": "chat", "completion_params": {}},144 "query": ["sys", "query"],145 "parameters": [{"name": "location", "type": "string", "description": "location", "required": True}],146 "instruction": "",147 "reasoning_mode": "function_call",148 "memory": None,149 },150 }151 )152 153 node._fetch_model_config = get_mocked_fetch_model_config(154 provider="openai",155 model="gpt-3.5-turbo",156 mode="chat",157 credentials={"openai_api_key": os.environ.get("OPENAI_API_KEY")},158 )159 db.session.close = MagicMock()160 161 # construct variable pool162 pool = VariablePool(163 system_variables={164 SystemVariableKey.QUERY: "what's the weather in SF",165 SystemVariableKey.FILES: [],166 SystemVariableKey.CONVERSATION_ID: "abababa",167 SystemVariableKey.USER_ID: "aaa",168 },169 user_inputs={},170 environment_variables=[],171 )172 173 result = node._run()174 175 assert result.status == WorkflowNodeExecutionStatus.SUCCEEDED176 assert result.outputs is not None177 assert result.outputs.get("location") == "kawaii"178 assert result.outputs.get("__reason") == None179 180 181@pytest.mark.parametrize("setup_openai_mock", [["chat"]], indirect=True)182def test_instructions(setup_openai_mock):183 """184 Test chat parameter extractor.185 """186 node = init_parameter_extractor_node(187 config={188 "id": "llm",189 "data": {190 "title": "123",191 "type": "parameter-extractor",192 "model": {"provider": "openai", "name": "gpt-3.5-turbo", "mode": "chat", "completion_params": {}},193 "query": ["sys", "query"],194 "parameters": [{"name": "location", "type": "string", "description": "location", "required": True}],195 "reasoning_mode": "function_call",196 "instruction": "{{#sys.query#}}",197 "memory": None,198 },199 },200 )201 202 node._fetch_model_config = get_mocked_fetch_model_config(203 provider="openai",204 model="gpt-3.5-turbo",205 mode="chat",206 credentials={"openai_api_key": os.environ.get("OPENAI_API_KEY")},207 )208 db.session.close = MagicMock()209 210 result = node._run()211 212 assert result.status == WorkflowNodeExecutionStatus.SUCCEEDED213 assert result.outputs is not None214 assert result.outputs.get("location") == "kawaii"215 assert result.outputs.get("__reason") == None216 217 process_data = result.process_data218 219 assert process_data is not None220 process_data.get("prompts")221 222 for prompt in process_data.get("prompts", []):223 if prompt.get("role") == "system":224 assert "what's the weather in SF" in prompt.get("text")225 226 227@pytest.mark.parametrize("setup_anthropic_mock", [["none"]], indirect=True)228def test_chat_parameter_extractor(setup_anthropic_mock):229 """230 Test chat parameter extractor.231 """232 node = init_parameter_extractor_node(233 config={234 "id": "llm",235 "data": {236 "title": "123",237 "type": "parameter-extractor",238 "model": {"provider": "anthropic", "name": "claude-2", "mode": "chat", "completion_params": {}},239 "query": ["sys", "query"],240 "parameters": [{"name": "location", "type": "string", "description": "location", "required": True}],241 "reasoning_mode": "prompt",242 "instruction": "",243 "memory": None,244 },245 },246 )247 248 node._fetch_model_config = get_mocked_fetch_model_config(249 provider="anthropic",250 model="claude-2",251 mode="chat",252 credentials={"anthropic_api_key": os.environ.get("ANTHROPIC_API_KEY")},253 )254 db.session.close = MagicMock()255 256 result = node._run()257 258 assert result.status == WorkflowNodeExecutionStatus.SUCCEEDED259 assert result.outputs is not None260 assert result.outputs.get("location") == ""261 assert (262 result.outputs.get("__reason")263 == "Failed to extract result from function call or text response, using empty result."264 )265 assert result.process_data is not None266 prompts = result.process_data.get("prompts", [])267 268 for prompt in prompts:269 if prompt.get("role") == "user":270 if "<structure>" in prompt.get("text"):271 assert '<structure>\n{"type": "object"' in prompt.get("text")272 273 274@pytest.mark.parametrize("setup_openai_mock", [["completion"]], indirect=True)275def test_completion_parameter_extractor(setup_openai_mock):276 """277 Test completion parameter extractor.278 """279 node = init_parameter_extractor_node(280 config={281 "id": "llm",282 "data": {283 "title": "123",284 "type": "parameter-extractor",285 "model": {286 "provider": "openai",287 "name": "gpt-3.5-turbo-instruct",288 "mode": "completion",289 "completion_params": {},290 },291 "query": ["sys", "query"],292 "parameters": [{"name": "location", "type": "string", "description": "location", "required": True}],293 "reasoning_mode": "prompt",294 "instruction": "{{#sys.query#}}",295 "memory": None,296 },297 },298 )299 300 node._fetch_model_config = get_mocked_fetch_model_config(301 provider="openai",302 model="gpt-3.5-turbo-instruct",303 mode="completion",304 credentials={"openai_api_key": os.environ.get("OPENAI_API_KEY")},305 )306 db.session.close = MagicMock()307 308 result = node._run()309 310 assert result.status == WorkflowNodeExecutionStatus.SUCCEEDED311 assert result.outputs is not None312 assert result.outputs.get("location") == ""313 assert (314 result.outputs.get("__reason")315 == "Failed to extract result from function call or text response, using empty result."316 )317 assert result.process_data is not None318 assert len(result.process_data.get("prompts", [])) == 1319 assert "SF" in result.process_data.get("prompts", [])[0].get("text")320 321 322def test_extract_json_response():323 """324 Test extract json response.325 """326 327 node = init_parameter_extractor_node(328 config={329 "id": "llm",330 "data": {331 "title": "123",332 "type": "parameter-extractor",333 "model": {334 "provider": "openai",335 "name": "gpt-3.5-turbo-instruct",336 "mode": "completion",337 "completion_params": {},338 },339 "query": ["sys", "query"],340 "parameters": [{"name": "location", "type": "string", "description": "location", "required": True}],341 "reasoning_mode": "prompt",342 "instruction": "{{#sys.query#}}",343 "memory": None,344 },345 },346 )347 348 result = node._extract_complete_json_response("""349 uwu{ovo}350 {351 "location": "kawaii"352 }353 hello world.354 """)355 356 assert result is not None357 assert result["location"] == "kawaii"358 359 360@pytest.mark.parametrize("setup_anthropic_mock", [["none"]], indirect=True)361def test_chat_parameter_extractor_with_memory(setup_anthropic_mock):362 """363 Test chat parameter extractor with memory.364 """365 node = init_parameter_extractor_node(366 config={367 "id": "llm",368 "data": {369 "title": "123",370 "type": "parameter-extractor",371 "model": {"provider": "anthropic", "name": "claude-2", "mode": "chat", "completion_params": {}},372 "query": ["sys", "query"],373 "parameters": [{"name": "location", "type": "string", "description": "location", "required": True}],374 "reasoning_mode": "prompt",375 "instruction": "",376 "memory": {"window": {"enabled": True, "size": 50}},377 },378 },379 )380 381 node._fetch_model_config = get_mocked_fetch_model_config(382 provider="anthropic",383 model="claude-2",384 mode="chat",385 credentials={"anthropic_api_key": os.environ.get("ANTHROPIC_API_KEY")},386 )387 node._fetch_memory = get_mocked_fetch_memory("customized memory")388 db.session.close = MagicMock()389 390 result = node._run()391 392 assert result.status == WorkflowNodeExecutionStatus.SUCCEEDED393 assert result.outputs is not None394 assert result.outputs.get("location") == ""395 assert (396 result.outputs.get("__reason")397 == "Failed to extract result from function call or text response, using empty result."398 )399 assert result.process_data is not None400 prompts = result.process_data.get("prompts", [])401 402 latest_role = None403 for prompt in prompts:404 if prompt.get("role") == "user":405 if "<structure>" in prompt.get("text"):406 assert '<structure>\n{"type": "object"' in prompt.get("text")407 elif prompt.get("role") == "system":408 assert "customized memory" in prompt.get("text")409 410 if latest_role is not None:411 assert latest_role != prompt.get("role")412 413 if prompt.get("role") in {"user", "assistant"}:414 latest_role = prompt.get("role")415 