Underground-Digital/Workflow-Engine
0
1import os
2from collections.abc import Generator
3
4import pytest
5
6from core.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta
7from core.model_runtime.entities.message_entities import (
8 AssistantPromptMessage,
9 PromptMessageTool,
10 SystemPromptMessage,
11 UserPromptMessage,
12)
13from core.model_runtime.entities.model_entities import AIModelEntity
14from core.model_runtime.errors.validate import CredentialsValidateFailedError
15from core.model_runtime.model_providers.gitee_ai.llm.llm import GiteeAILargeLanguageModel
16
17
18def test_predefined_models():
19 model = GiteeAILargeLanguageModel()
20 model_schemas = model.predefined_models()
21
22 assert len(model_schemas) >= 1
23 assert isinstance(model_schemas[0], AIModelEntity)
24
25
26def test_validate_credentials_for_chat_model():
27 model = GiteeAILargeLanguageModel()
28
29 with pytest.raises(CredentialsValidateFailedError):
30 # model name to gpt-3.5-turbo because of mocking
31 model.validate_credentials(model="gpt-3.5-turbo", credentials={"api_key": "invalid_key"})
32
33 model.validate_credentials(
34 model="Qwen2-7B-Instruct",
35 credentials={"api_key": os.environ.get("GITEE_AI_API_KEY")},
36 )
37
38
39def test_invoke_chat_model():
40 model = GiteeAILargeLanguageModel()
41
42 result = model.invoke(
43 model="Qwen2-7B-Instruct",
44 credentials={"api_key": os.environ.get("GITEE_AI_API_KEY")},
45 prompt_messages=[
46 SystemPromptMessage(
47 content="You are a helpful AI assistant.",
48 ),
49 UserPromptMessage(content="Hello World!"),
50 ],
51 model_parameters={
52 "temperature": 0.0,
53 "top_p": 1.0,
54 "presence_penalty": 0.0,
55 "frequency_penalty": 0.0,
56 "max_tokens": 10,
57 "stream": False,
58 },
59 stop=["How"],
60 stream=False,
61 user="foo",
62 )
63
64 assert isinstance(result, LLMResult)
65 assert len(result.message.content) > 0
66
67
68def test_invoke_stream_chat_model():
69 model = GiteeAILargeLanguageModel()
70
71 result = model.invoke(
72 model="Qwen2-7B-Instruct",
73 credentials={"api_key": os.environ.get("GITEE_AI_API_KEY")},
74 prompt_messages=[
75 SystemPromptMessage(
76 content="You are a helpful AI assistant.",
77 ),
78 UserPromptMessage(content="Hello World!"),
79 ],
80 model_parameters={"temperature": 0.0, "max_tokens": 100, "stream": False},
81 stream=True,
82 user="foo",
83 )
84
85 assert isinstance(result, Generator)
86
87 for chunk in result:
88 assert isinstance(chunk, LLMResultChunk)
89 assert isinstance(chunk.delta, LLMResultChunkDelta)
90 assert isinstance(chunk.delta.message, AssistantPromptMessage)
91 assert len(chunk.delta.message.content) > 0 if chunk.delta.finish_reason is None else True
92 if chunk.delta.finish_reason is not None:
93 assert chunk.delta.usage is not None
94
95
96def test_get_num_tokens():
97 model = GiteeAILargeLanguageModel()
98
99 num_tokens = model.get_num_tokens(
100 model="Qwen2-7B-Instruct",
101 credentials={"api_key": os.environ.get("GITEE_AI_API_KEY")},
102 prompt_messages=[UserPromptMessage(content="Hello World!")],
103 )
104
105 assert num_tokens == 10
106
107 num_tokens = model.get_num_tokens(
108 model="Qwen2-7B-Instruct",
109 credentials={"api_key": os.environ.get("GITEE_AI_API_KEY")},
110 prompt_messages=[
111 SystemPromptMessage(
112 content="You are a helpful AI assistant.",
113 ),
114 UserPromptMessage(content="Hello World!"),
115 ],
116 tools=[
117 PromptMessageTool(
118 name="get_weather",
119 description="Determine weather in my location",
120 parameters={
121 "type": "object",
122 "properties": {
123 "location": {"type": "string", "description": "The city and state e.g. San Francisco, CA"},
124 "unit": {"type": "string", "enum": ["c", "f"]},
125 },
126 "required": ["location"],
127 },
128 ),
129 ],
130 )
131
132 assert num_tokens == 77
133 