Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
test_llm.py230 linesDownload Raw Back to chatglm
1import os2from collections.abc import Generator3 4import pytest5 6from core.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta7from core.model_runtime.entities.message_entities import (8    AssistantPromptMessage,9    PromptMessageTool,10    SystemPromptMessage,11    UserPromptMessage,12)13from core.model_runtime.entities.model_entities import AIModelEntity14from core.model_runtime.errors.validate import CredentialsValidateFailedError15from core.model_runtime.model_providers.chatglm.llm.llm import ChatGLMLargeLanguageModel16from tests.integration_tests.model_runtime.__mock.openai import setup_openai_mock17 18 19def test_predefined_models():20    model = ChatGLMLargeLanguageModel()21    model_schemas = model.predefined_models()22    assert len(model_schemas) >= 123    assert isinstance(model_schemas[0], AIModelEntity)24 25 26@pytest.mark.parametrize("setup_openai_mock", [["chat"]], indirect=True)27def test_validate_credentials_for_chat_model(setup_openai_mock):28    model = ChatGLMLargeLanguageModel()29 30    with pytest.raises(CredentialsValidateFailedError):31        model.validate_credentials(model="chatglm2-6b", credentials={"api_base": "invalid_key"})32 33    model.validate_credentials(model="chatglm2-6b", credentials={"api_base": os.environ.get("CHATGLM_API_BASE")})34 35 36@pytest.mark.parametrize("setup_openai_mock", [["chat"]], indirect=True)37def test_invoke_model(setup_openai_mock):38    model = ChatGLMLargeLanguageModel()39 40    response = model.invoke(41        model="chatglm2-6b",42        credentials={"api_base": os.environ.get("CHATGLM_API_BASE")},43        prompt_messages=[44            SystemPromptMessage(45                content="You are a helpful AI assistant.",46            ),47            UserPromptMessage(content="Hello World!"),48        ],49        model_parameters={50            "temperature": 0.7,51            "top_p": 1.0,52        },53        stop=["you"],54        user="abc-123",55        stream=False,56    )57 58    assert isinstance(response, LLMResult)59    assert len(response.message.content) > 060    assert response.usage.total_tokens > 061 62 63@pytest.mark.parametrize("setup_openai_mock", [["chat"]], indirect=True)64def test_invoke_stream_model(setup_openai_mock):65    model = ChatGLMLargeLanguageModel()66 67    response = model.invoke(68        model="chatglm2-6b",69        credentials={"api_base": os.environ.get("CHATGLM_API_BASE")},70        prompt_messages=[71            SystemPromptMessage(72                content="You are a helpful AI assistant.",73            ),74            UserPromptMessage(content="Hello World!"),75        ],76        model_parameters={77            "temperature": 0.7,78            "top_p": 1.0,79        },80        stop=["you"],81        stream=True,82        user="abc-123",83    )84 85    assert isinstance(response, Generator)86    for chunk in response:87        assert isinstance(chunk, LLMResultChunk)88        assert isinstance(chunk.delta, LLMResultChunkDelta)89        assert isinstance(chunk.delta.message, AssistantPromptMessage)90        assert len(chunk.delta.message.content) > 0 if chunk.delta.finish_reason is None else True91 92 93@pytest.mark.parametrize("setup_openai_mock", [["chat"]], indirect=True)94def test_invoke_stream_model_with_functions(setup_openai_mock):95    model = ChatGLMLargeLanguageModel()96 97    response = model.invoke(98        model="chatglm3-6b",99        credentials={"api_base": os.environ.get("CHATGLM_API_BASE")},100        prompt_messages=[101            SystemPromptMessage(102                content="你是一个天气机器人,你不知道今天的天气怎么样,你需要通过调用一个函数来获取天气信息。"103            ),104            UserPromptMessage(content="波士顿天气如何?"),105        ],106        model_parameters={107            "temperature": 0,108            "top_p": 1.0,109        },110        stop=["you"],111        user="abc-123",112        stream=True,113        tools=[114            PromptMessageTool(115                name="get_current_weather",116                description="Get the current weather in a given location",117                parameters={118                    "type": "object",119                    "properties": {120                        "location": {"type": "string", "description": "The city and state e.g. San Francisco, CA"},121                        "unit": {"type": "string", "enum": ["celsius", "fahrenheit"]},122                    },123                    "required": ["location"],124                },125            )126        ],127    )128 129    assert isinstance(response, Generator)130 131    call: LLMResultChunk = None132    chunks = []133 134    for chunk in response:135        chunks.append(chunk)136        assert isinstance(chunk, LLMResultChunk)137        assert isinstance(chunk.delta, LLMResultChunkDelta)138        assert isinstance(chunk.delta.message, AssistantPromptMessage)139        assert len(chunk.delta.message.content) > 0 if chunk.delta.finish_reason is None else True140 141        if chunk.delta.message.tool_calls and len(chunk.delta.message.tool_calls) > 0:142            call = chunk143            break144 145    assert call is not None146    assert call.delta.message.tool_calls[0].function.name == "get_current_weather"147 148 149@pytest.mark.parametrize("setup_openai_mock", [["chat"]], indirect=True)150def test_invoke_model_with_functions(setup_openai_mock):151    model = ChatGLMLargeLanguageModel()152 153    response = model.invoke(154        model="chatglm3-6b",155        credentials={"api_base": os.environ.get("CHATGLM_API_BASE")},156        prompt_messages=[UserPromptMessage(content="What is the weather like in San Francisco?")],157        model_parameters={158            "temperature": 0.7,159            "top_p": 1.0,160        },161        stop=["you"],162        user="abc-123",163        stream=False,164        tools=[165            PromptMessageTool(166                name="get_current_weather",167                description="Get the current weather in a given location",168                parameters={169                    "type": "object",170                    "properties": {171                        "location": {"type": "string", "description": "The city and state e.g. San Francisco, CA"},172                        "unit": {"type": "string", "enum": ["c", "f"]},173                    },174                    "required": ["location"],175                },176            )177        ],178    )179 180    assert isinstance(response, LLMResult)181    assert len(response.message.content) > 0182    assert response.usage.total_tokens > 0183    assert response.message.tool_calls[0].function.name == "get_current_weather"184 185 186def test_get_num_tokens():187    model = ChatGLMLargeLanguageModel()188 189    num_tokens = model.get_num_tokens(190        model="chatglm2-6b",191        credentials={"api_base": os.environ.get("CHATGLM_API_BASE")},192        prompt_messages=[193            SystemPromptMessage(194                content="You are a helpful AI assistant.",195            ),196            UserPromptMessage(content="Hello World!"),197        ],198        tools=[199            PromptMessageTool(200                name="get_current_weather",201                description="Get the current weather in a given location",202                parameters={203                    "type": "object",204                    "properties": {205                        "location": {"type": "string", "description": "The city and state e.g. San Francisco, CA"},206                        "unit": {"type": "string", "enum": ["c", "f"]},207                    },208                    "required": ["location"],209                },210            )211        ],212    )213 214    assert isinstance(num_tokens, int)215    assert num_tokens == 77216 217    num_tokens = model.get_num_tokens(218        model="chatglm2-6b",219        credentials={"api_base": os.environ.get("CHATGLM_API_BASE")},220        prompt_messages=[221            SystemPromptMessage(222                content="You are a helpful AI assistant.",223            ),224            UserPromptMessage(content="Hello World!"),225        ],226    )227 228    assert isinstance(num_tokens, int)229    assert num_tokens == 21230