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.huggingface_hub.text_embedding.text_embedding import (8 HuggingfaceHubTextEmbeddingModel,9)10 11 12def test_hosted_inference_api_validate_credentials():13 model = HuggingfaceHubTextEmbeddingModel()14 15 with pytest.raises(CredentialsValidateFailedError):16 model.validate_credentials(17 model="facebook/bart-base",18 credentials={19 "huggingfacehub_api_type": "hosted_inference_api",20 "huggingfacehub_api_token": "invalid_key",21 },22 )23 24 model.validate_credentials(25 model="facebook/bart-base",26 credentials={27 "huggingfacehub_api_type": "hosted_inference_api",28 "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),29 },30 )31 32 33def test_hosted_inference_api_invoke_model():34 model = HuggingfaceHubTextEmbeddingModel()35 36 result = model.invoke(37 model="facebook/bart-base",38 credentials={39 "huggingfacehub_api_type": "hosted_inference_api",40 "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),41 },42 texts=["hello", "world"],43 )44 45 assert isinstance(result, TextEmbeddingResult)46 assert len(result.embeddings) == 247 assert result.usage.total_tokens == 248 49 50def test_inference_endpoints_validate_credentials():51 model = HuggingfaceHubTextEmbeddingModel()52 53 with pytest.raises(CredentialsValidateFailedError):54 model.validate_credentials(55 model="all-MiniLM-L6-v2",56 credentials={57 "huggingfacehub_api_type": "inference_endpoints",58 "huggingfacehub_api_token": "invalid_key",59 "huggingface_namespace": "Dify-AI",60 "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_EMBEDDINGS_ENDPOINT_URL"),61 "task_type": "feature-extraction",62 },63 )64 65 model.validate_credentials(66 model="all-MiniLM-L6-v2",67 credentials={68 "huggingfacehub_api_type": "inference_endpoints",69 "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),70 "huggingface_namespace": "Dify-AI",71 "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_EMBEDDINGS_ENDPOINT_URL"),72 "task_type": "feature-extraction",73 },74 )75 76 77def test_inference_endpoints_invoke_model():78 model = HuggingfaceHubTextEmbeddingModel()79 80 result = model.invoke(81 model="all-MiniLM-L6-v2",82 credentials={83 "huggingfacehub_api_type": "inference_endpoints",84 "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),85 "huggingface_namespace": "Dify-AI",86 "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_EMBEDDINGS_ENDPOINT_URL"),87 "task_type": "feature-extraction",88 },89 texts=["hello", "world"],90 )91 92 assert isinstance(result, TextEmbeddingResult)93 assert len(result.embeddings) == 294 assert result.usage.total_tokens == 095 96 97def test_get_num_tokens():98 model = HuggingfaceHubTextEmbeddingModel()99 100 num_tokens = model.get_num_tokens(101 model="all-MiniLM-L6-v2",102 credentials={103 "huggingfacehub_api_type": "inference_endpoints",104 "huggingfacehub_api_token": os.environ.get("HUGGINGFACE_API_KEY"),105 "huggingface_namespace": "Dify-AI",106 "huggingfacehub_endpoint_url": os.environ.get("HUGGINGFACE_EMBEDDINGS_ENDPOINT_URL"),107 "task_type": "feature-extraction",108 },109 texts=["hello", "world"],110 )111 112 assert num_tokens == 2113 