Underground-Digital/Workflow-Engine
0
1import os2 3import pytest4 5from core.model_runtime.entities.text_embedding_entities import TextEmbeddingResult6from core.model_runtime.errors.validate import CredentialsValidateFailedError7from core.model_runtime.model_providers.xinference.text_embedding.text_embedding import XinferenceTextEmbeddingModel8from tests.integration_tests.model_runtime.__mock.xinference import MOCK, setup_xinference_mock9 10 11@pytest.mark.parametrize("setup_xinference_mock", [["none"]], indirect=True)12def test_validate_credentials(setup_xinference_mock):13 model = XinferenceTextEmbeddingModel()14 15 with pytest.raises(CredentialsValidateFailedError):16 model.validate_credentials(17 model="bge-base-en",18 credentials={19 "server_url": os.environ.get("XINFERENCE_SERVER_URL"),20 "model_uid": "www " + os.environ.get("XINFERENCE_EMBEDDINGS_MODEL_UID"),21 },22 )23 24 model.validate_credentials(25 model="bge-base-en",26 credentials={27 "server_url": os.environ.get("XINFERENCE_SERVER_URL"),28 "model_uid": os.environ.get("XINFERENCE_EMBEDDINGS_MODEL_UID"),29 },30 )31 32 33@pytest.mark.parametrize("setup_xinference_mock", [["none"]], indirect=True)34def test_invoke_model(setup_xinference_mock):35 model = XinferenceTextEmbeddingModel()36 37 result = model.invoke(38 model="bge-base-en",39 credentials={40 "server_url": os.environ.get("XINFERENCE_SERVER_URL"),41 "model_uid": os.environ.get("XINFERENCE_EMBEDDINGS_MODEL_UID"),42 },43 texts=["hello", "world"],44 user="abc-123",45 )46 47 assert isinstance(result, TextEmbeddingResult)48 assert len(result.embeddings) == 249 assert result.usage.total_tokens > 050 51 52def test_get_num_tokens():53 model = XinferenceTextEmbeddingModel()54 55 num_tokens = model.get_num_tokens(56 model="bge-base-en",57 credentials={58 "server_url": os.environ.get("XINFERENCE_SERVER_URL"),59 "model_uid": os.environ.get("XINFERENCE_EMBEDDINGS_MODEL_UID"),60 },61 texts=["hello", "world"],62 )63 64 assert num_tokens == 265 