Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
test_llm.py133 linesDownload Raw Back to gitee_ai
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