Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
test_parameter_extractor.py415 linesDownload Raw Back to nodes
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