Underground-Digital/Workflow-Engine
0
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 