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 AssistantPromptMessage, UserPromptMessage8from core.model_runtime.errors.validate import CredentialsValidateFailedError9from core.model_runtime.model_providers.huggingface_hub.llm.llm import HuggingfaceHubLargeLanguageModel10from tests.integration_tests.model_runtime.__mock.huggingface import setup_huggingface_mock11 12 13@pytest.mark.parametrize("setup_huggingface_mock", [["none"]], indirect=True)14def test_hosted_inference_api_validate_credentials(setup_huggingface_mock):15 model = HuggingfaceHubLargeLanguageModel()16 17 with pytest.raises(CredentialsValidateFailedError):18 model.validate_credentials(19 model="HuggingFaceH4/zephyr-7b-beta",20 credentials={"huggingfacehub_api_type": "hosted_inference_api", "huggingfacehub_api_token": "invalid_key"},21 )22 23 with pytest.raises(CredentialsValidateFailedError):24 model.validate_credentials(25 model="fake-model",26 credentials={"huggingfacehub_api_type": "hosted_inference_api", "huggingfacehub_api_token": "invalid_key"},27 )28 29 model.validate_credentials(30 model="HuggingFaceH4/zephyr-7b-beta",31 credentials={32 "huggingfacehub_api_type": "hosted_inference_api",33 "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),34 },35 )36 37 38@pytest.mark.parametrize("setup_huggingface_mock", [["none"]], indirect=True)39def test_hosted_inference_api_invoke_model(setup_huggingface_mock):40 model = HuggingfaceHubLargeLanguageModel()41 42 response = model.invoke(43 model="HuggingFaceH4/zephyr-7b-beta",44 credentials={45 "huggingfacehub_api_type": "hosted_inference_api",46 "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),47 },48 prompt_messages=[UserPromptMessage(content="Who are you?")],49 model_parameters={50 "temperature": 1.0,51 "top_k": 2,52 "top_p": 0.5,53 },54 stop=["How"],55 stream=False,56 user="abc-123",57 )58 59 assert isinstance(response, LLMResult)60 assert len(response.message.content) > 061 62 63@pytest.mark.parametrize("setup_huggingface_mock", [["none"]], indirect=True)64def test_hosted_inference_api_invoke_stream_model(setup_huggingface_mock):65 model = HuggingfaceHubLargeLanguageModel()66 67 response = model.invoke(68 model="HuggingFaceH4/zephyr-7b-beta",69 credentials={70 "huggingfacehub_api_type": "hosted_inference_api",71 "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),72 },73 prompt_messages=[UserPromptMessage(content="Who are you?")],74 model_parameters={75 "temperature": 1.0,76 "top_k": 2,77 "top_p": 0.5,78 },79 stop=["How"],80 stream=True,81 user="abc-123",82 )83 84 assert isinstance(response, Generator)85 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_huggingface_mock", [["none"]], indirect=True)94def test_inference_endpoints_text_generation_validate_credentials(setup_huggingface_mock):95 model = HuggingfaceHubLargeLanguageModel()96 97 with pytest.raises(CredentialsValidateFailedError):98 model.validate_credentials(99 model="openchat/openchat_3.5",100 credentials={101 "huggingfacehub_api_type": "inference_endpoints",102 "huggingfacehub_api_token": "invalid_key",103 "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_TEXT_GEN_ENDPOINT_URL"),104 "task_type": "text-generation",105 },106 )107 108 model.validate_credentials(109 model="openchat/openchat_3.5",110 credentials={111 "huggingfacehub_api_type": "inference_endpoints",112 "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),113 "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_TEXT_GEN_ENDPOINT_URL"),114 "task_type": "text-generation",115 },116 )117 118 119@pytest.mark.parametrize("setup_huggingface_mock", [["none"]], indirect=True)120def test_inference_endpoints_text_generation_invoke_model(setup_huggingface_mock):121 model = HuggingfaceHubLargeLanguageModel()122 123 response = model.invoke(124 model="openchat/openchat_3.5",125 credentials={126 "huggingfacehub_api_type": "inference_endpoints",127 "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),128 "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_TEXT_GEN_ENDPOINT_URL"),129 "task_type": "text-generation",130 },131 prompt_messages=[UserPromptMessage(content="Who are you?")],132 model_parameters={133 "temperature": 1.0,134 "top_k": 2,135 "top_p": 0.5,136 },137 stop=["How"],138 stream=False,139 user="abc-123",140 )141 142 assert isinstance(response, LLMResult)143 assert len(response.message.content) > 0144 145 146@pytest.mark.parametrize("setup_huggingface_mock", [["none"]], indirect=True)147def test_inference_endpoints_text_generation_invoke_stream_model(setup_huggingface_mock):148 model = HuggingfaceHubLargeLanguageModel()149 150 response = model.invoke(151 model="openchat/openchat_3.5",152 credentials={153 "huggingfacehub_api_type": "inference_endpoints",154 "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),155 "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_TEXT_GEN_ENDPOINT_URL"),156 "task_type": "text-generation",157 },158 prompt_messages=[UserPromptMessage(content="Who are you?")],159 model_parameters={160 "temperature": 1.0,161 "top_k": 2,162 "top_p": 0.5,163 },164 stop=["How"],165 stream=True,166 user="abc-123",167 )168 169 assert isinstance(response, Generator)170 171 for chunk in response:172 assert isinstance(chunk, LLMResultChunk)173 assert isinstance(chunk.delta, LLMResultChunkDelta)174 assert isinstance(chunk.delta.message, AssistantPromptMessage)175 assert len(chunk.delta.message.content) > 0 if chunk.delta.finish_reason is None else True176 177 178@pytest.mark.parametrize("setup_huggingface_mock", [["none"]], indirect=True)179def test_inference_endpoints_text2text_generation_validate_credentials(setup_huggingface_mock):180 model = HuggingfaceHubLargeLanguageModel()181 182 with pytest.raises(CredentialsValidateFailedError):183 model.validate_credentials(184 model="google/mt5-base",185 credentials={186 "huggingfacehub_api_type": "inference_endpoints",187 "huggingfacehub_api_token": "invalid_key",188 "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_TEXT2TEXT_GEN_ENDPOINT_URL"),189 "task_type": "text2text-generation",190 },191 )192 193 model.validate_credentials(194 model="google/mt5-base",195 credentials={196 "huggingfacehub_api_type": "inference_endpoints",197 "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),198 "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_TEXT2TEXT_GEN_ENDPOINT_URL"),199 "task_type": "text2text-generation",200 },201 )202 203 204@pytest.mark.parametrize("setup_huggingface_mock", [["none"]], indirect=True)205def test_inference_endpoints_text2text_generation_invoke_model(setup_huggingface_mock):206 model = HuggingfaceHubLargeLanguageModel()207 208 response = model.invoke(209 model="google/mt5-base",210 credentials={211 "huggingfacehub_api_type": "inference_endpoints",212 "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),213 "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_TEXT2TEXT_GEN_ENDPOINT_URL"),214 "task_type": "text2text-generation",215 },216 prompt_messages=[UserPromptMessage(content="Who are you?")],217 model_parameters={218 "temperature": 1.0,219 "top_k": 2,220 "top_p": 0.5,221 },222 stop=["How"],223 stream=False,224 user="abc-123",225 )226 227 assert isinstance(response, LLMResult)228 assert len(response.message.content) > 0229 230 231@pytest.mark.parametrize("setup_huggingface_mock", [["none"]], indirect=True)232def test_inference_endpoints_text2text_generation_invoke_stream_model(setup_huggingface_mock):233 model = HuggingfaceHubLargeLanguageModel()234 235 response = model.invoke(236 model="google/mt5-base",237 credentials={238 "huggingfacehub_api_type": "inference_endpoints",239 "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),240 "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_TEXT2TEXT_GEN_ENDPOINT_URL"),241 "task_type": "text2text-generation",242 },243 prompt_messages=[UserPromptMessage(content="Who are you?")],244 model_parameters={245 "temperature": 1.0,246 "top_k": 2,247 "top_p": 0.5,248 },249 stop=["How"],250 stream=True,251 user="abc-123",252 )253 254 assert isinstance(response, Generator)255 256 for chunk in response:257 assert isinstance(chunk, LLMResultChunk)258 assert isinstance(chunk.delta, LLMResultChunkDelta)259 assert isinstance(chunk.delta.message, AssistantPromptMessage)260 assert len(chunk.delta.message.content) > 0 if chunk.delta.finish_reason is None else True261 262 263def test_get_num_tokens():264 model = HuggingfaceHubLargeLanguageModel()265 266 num_tokens = model.get_num_tokens(267 model="google/mt5-base",268 credentials={269 "huggingfacehub_api_type": "inference_endpoints",270 "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),271 "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_TEXT2TEXT_GEN_ENDPOINT_URL"),272 "task_type": "text2text-generation",273 },274 prompt_messages=[UserPromptMessage(content="Hello World!")],275 )276 277 assert num_tokens == 7278 