-
Notifications
You must be signed in to change notification settings - Fork 16
Expand file tree
/
Copy pathtest_assistant.py
More file actions
135 lines (115 loc) · 4.62 KB
/
Copy pathtest_assistant.py
File metadata and controls
135 lines (115 loc) · 4.62 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
"""Provider adapter tests for the local LLM assistant."""
import unittest
from unittest.mock import Mock, patch
import assistant
import config
class TestOpenAICompatibleProvider(unittest.TestCase):
def setUp(self):
self._config = patch.multiple(
config,
ASSISTANT_PROVIDER="openai",
OPENAI_URL="http://localhost:8080/v1/",
OPENAI_MODEL="local-test-model",
OPENAI_API_KEY="",
)
self._config.start()
def tearDown(self):
self._config.stop()
@patch.object(assistant, "_system_prompt", return_value="system prompt")
@patch.object(assistant.requests, "post")
def test_chat_completion_tool_call_is_normalized(self, post, _prompt):
response = Mock()
response.json.return_value = {
"choices": [{
"message": {
"tool_calls": [{
"function": {
"name": "save_note",
"arguments": '{"content":"buy milk"}',
},
}],
},
}],
}
post.return_value = response
result = assistant._call_openai("remember to buy milk")
self.assertEqual(result, {
"function": "save_note",
"arguments": {"content": "buy milk"},
})
response.raise_for_status.assert_called_once_with()
url = post.call_args.args[0]
kwargs = post.call_args.kwargs
self.assertEqual(url, "http://localhost:8080/v1/chat/completions")
self.assertEqual(kwargs["json"]["model"], "local-test-model")
self.assertEqual(kwargs["json"]["messages"][0]["content"],
"system prompt")
self.assertEqual(kwargs["json"]["messages"][1], {
"role": "user",
"content": "remember to buy milk",
})
self.assertIs(kwargs["json"]["tools"], assistant.TOOLS)
self.assertNotIn("Authorization", kwargs["headers"])
self.assertEqual(kwargs["timeout"], 120)
@patch.object(assistant.requests, "post")
def test_invalid_tool_arguments_are_rejected(self, post):
response = Mock()
response.json.return_value = {
"choices": [{
"message": {
"tool_calls": [{
"function": {
"name": "save_note",
"arguments": "not json",
},
}],
},
}],
}
post.return_value = response
self.assertIsNone(assistant._call_openai("remember this"))
@patch.object(assistant.requests, "post")
def test_object_tool_arguments_are_also_accepted(self, post):
response = Mock()
response.json.return_value = {
"choices": [{
"message": {
"tool_calls": [{
"function": {
"name": "save_note",
"arguments": {"content": "already decoded"},
},
}],
},
}],
}
post.return_value = response
self.assertEqual(assistant._call_openai("remember this"), {
"function": "save_note",
"arguments": {"content": "already decoded"},
})
@patch.object(assistant.requests, "get")
def test_health_check_uses_models_endpoint_and_api_key(self, get):
config.OPENAI_API_KEY = "local-secret"
get.return_value.status_code = 200
self.assertTrue(assistant.ping_provider())
get.assert_called_once_with(
"http://localhost:8080/v1/models",
headers={"Authorization": "Bearer local-secret"},
timeout=2,
)
class TestProviderDispatch(unittest.TestCase):
@patch.object(assistant, "_call_openai", return_value={"provider": "openai"})
def test_openai_provider_is_selected(self, call_openai):
with patch.object(config, "ASSISTANT_PROVIDER", "openai"):
self.assertEqual(assistant._call_provider("hello"),
{"provider": "openai"})
call_openai.assert_called_once_with("hello")
@patch.object(assistant, "_call_ollama", return_value={"provider": "ollama"})
def test_ollama_remains_the_default(self, call_ollama):
with patch.object(config, "ASSISTANT_PROVIDER", "ollama"):
self.assertEqual(assistant._call_provider("hello"),
{"provider": "ollama"})
call_ollama.assert_called_once_with("hello")
if __name__ == "__main__":
unittest.main()