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.stepfun.llm.llm import StepfunLargeLanguageModel16 17 18def test_validate_credentials():19 model = StepfunLargeLanguageModel()20 21 with pytest.raises(CredentialsValidateFailedError):22 model.validate_credentials(model="step-1-8k", credentials={"api_key": "invalid_key"})23 24 model.validate_credentials(model="step-1-8k", credentials={"api_key": os.environ.get("STEPFUN_API_KEY")})25 26 27def test_invoke_model():28 model = StepfunLargeLanguageModel()29 30 response = model.invoke(31 model="step-1-8k",32 credentials={"api_key": os.environ.get("STEPFUN_API_KEY")},33 prompt_messages=[UserPromptMessage(content="Hello World!")],34 model_parameters={"temperature": 0.9, "top_p": 0.7},35 stop=["Hi"],36 stream=False,37 user="abc-123",38 )39 40 assert isinstance(response, LLMResult)41 assert len(response.message.content) > 042 43 44def test_invoke_stream_model():45 model = StepfunLargeLanguageModel()46 47 response = model.invoke(48 model="step-1-8k",49 credentials={"api_key": os.environ.get("STEPFUN_API_KEY")},50 prompt_messages=[51 SystemPromptMessage(52 content="You are a helpful AI assistant.",53 ),54 UserPromptMessage(content="Hello World!"),55 ],56 model_parameters={"temperature": 0.9, "top_p": 0.7},57 stream=True,58 user="abc-123",59 )60 61 assert isinstance(response, Generator)62 63 for chunk in response:64 assert isinstance(chunk, LLMResultChunk)65 assert isinstance(chunk.delta, LLMResultChunkDelta)66 assert isinstance(chunk.delta.message, AssistantPromptMessage)67 assert len(chunk.delta.message.content) > 0 if chunk.delta.finish_reason is None else True68 69 70def test_get_customizable_model_schema():71 model = StepfunLargeLanguageModel()72 73 schema = model.get_customizable_model_schema(74 model="step-1-8k", credentials={"api_key": os.environ.get("STEPFUN_API_KEY")}75 )76 assert isinstance(schema, AIModelEntity)77 78 79def test_invoke_chat_model_with_tools():80 model = StepfunLargeLanguageModel()81 82 result = model.invoke(83 model="step-1-8k",84 credentials={"api_key": os.environ.get("STEPFUN_API_KEY")},85 prompt_messages=[86 SystemPromptMessage(87 content="You are a helpful AI assistant.",88 ),89 UserPromptMessage(90 content="what's the weather today in Shanghai?",91 ),92 ],93 model_parameters={"temperature": 0.9, "max_tokens": 100},94 tools=[95 PromptMessageTool(96 name="get_weather",97 description="Determine weather in my location",98 parameters={99 "type": "object",100 "properties": {101 "location": {"type": "string", "description": "The city and state e.g. San Francisco, CA"},102 "unit": {"type": "string", "enum": ["c", "f"]},103 },104 "required": ["location"],105 },106 ),107 PromptMessageTool(108 name="get_stock_price",109 description="Get the current stock price",110 parameters={111 "type": "object",112 "properties": {"symbol": {"type": "string", "description": "The stock symbol"}},113 "required": ["symbol"],114 },115 ),116 ],117 stream=False,118 user="abc-123",119 )120 121 assert isinstance(result, LLMResult)122 assert isinstance(result.message, AssistantPromptMessage)123 assert len(result.message.tool_calls) > 0124 