rasti-ai 3.0.1 → 3.2.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- checksums.yaml +4 -4
- data/.features/api-normalization/api-design.md +184 -0
- data/.features/api-normalization/overview.md +230 -0
- data/AGENTS.md +167 -10
- data/README.md +161 -23
- data/lib/rasti/ai/anthropic/assistant.rb +1 -11
- data/lib/rasti/ai/anthropic/client.rb +6 -28
- data/lib/rasti/ai/anthropic/provider.rb +87 -0
- data/lib/rasti/ai/assistant.rb +10 -11
- data/lib/rasti/ai/client.rb +13 -10
- data/lib/rasti/ai/errors.rb +8 -0
- data/lib/rasti/ai/gemini/assistant.rb +1 -11
- data/lib/rasti/ai/gemini/client.rb +2 -24
- data/lib/rasti/ai/gemini/provider.rb +90 -0
- data/lib/rasti/ai/huawei_maas/assistant.rb +8 -0
- data/lib/rasti/ai/huawei_maas/client.rb +8 -0
- data/lib/rasti/ai/huawei_maas/provider.rb +25 -0
- data/lib/rasti/ai/huawei_maas/roles.rb +7 -0
- data/lib/rasti/ai/open_ai/assistant.rb +0 -4
- data/lib/rasti/ai/open_ai/client.rb +4 -26
- data/lib/rasti/ai/open_ai/provider.rb +74 -0
- data/lib/rasti/ai/open_router/assistant.rb +8 -0
- data/lib/rasti/ai/open_router/client.rb +8 -0
- data/lib/rasti/ai/open_router/provider.rb +25 -0
- data/lib/rasti/ai/open_router/roles.rb +7 -0
- data/lib/rasti/ai/provider.rb +197 -0
- data/lib/rasti/ai/provider_aware.rb +17 -0
- data/lib/rasti/ai/result.rb +11 -0
- data/lib/rasti/ai/roles.rb +11 -0
- data/lib/rasti/ai/version.rb +1 -1
- data/lib/rasti/ai.rb +101 -1
- data/spec/anthropic/provider_spec.rb +106 -0
- data/spec/gemini/provider_spec.rb +90 -0
- data/spec/huawei_maas/assistant_spec.rb +80 -0
- data/spec/huawei_maas/client_spec.rb +60 -0
- data/spec/minitest_helper.rb +6 -0
- data/spec/open_ai/provider_spec.rb +101 -0
- data/spec/open_router/assistant_spec.rb +80 -0
- data/spec/open_router/client_spec.rb +60 -0
- data/spec/rasti_ai_spec.rb +322 -0
- data/spec/resources/open_ai/conversation_request.json +1 -0
- data/spec/resources/open_ai/system_request.json +1 -0
- data/tasks/assistant.rake +5 -3
- metadata +39 -2
|
@@ -0,0 +1,197 @@
|
|
|
1
|
+
module Rasti
|
|
2
|
+
module AI
|
|
3
|
+
class Provider
|
|
4
|
+
|
|
5
|
+
MODULES = {
|
|
6
|
+
open_ai: 'OpenAI',
|
|
7
|
+
gemini: 'Gemini',
|
|
8
|
+
anthropic: 'Anthropic',
|
|
9
|
+
open_router: 'OpenRouter',
|
|
10
|
+
huawei_maas: 'HuaweiMaaS'
|
|
11
|
+
}.freeze
|
|
12
|
+
|
|
13
|
+
ALIASES = {
|
|
14
|
+
openai: :open_ai,
|
|
15
|
+
openrouter: :open_router,
|
|
16
|
+
huawei: :huawei_maas
|
|
17
|
+
}.freeze
|
|
18
|
+
|
|
19
|
+
class << self
|
|
20
|
+
|
|
21
|
+
def build(name, **options)
|
|
22
|
+
return name if name.is_a? Provider
|
|
23
|
+
provider_class(name).new(**options)
|
|
24
|
+
end
|
|
25
|
+
|
|
26
|
+
def names
|
|
27
|
+
MODULES.keys
|
|
28
|
+
end
|
|
29
|
+
|
|
30
|
+
private
|
|
31
|
+
|
|
32
|
+
def provider_class(name)
|
|
33
|
+
key = normalize name
|
|
34
|
+
key = ALIASES.fetch key, key
|
|
35
|
+
|
|
36
|
+
raise Errors::UnknownProvider.new(name, names) unless MODULES.key? key
|
|
37
|
+
|
|
38
|
+
AI.const_get(MODULES[key])::Provider
|
|
39
|
+
end
|
|
40
|
+
|
|
41
|
+
def normalize(name)
|
|
42
|
+
name.to_s.downcase.to_sym
|
|
43
|
+
end
|
|
44
|
+
|
|
45
|
+
end
|
|
46
|
+
|
|
47
|
+
attr_reader :model
|
|
48
|
+
|
|
49
|
+
def initialize(model:nil, api_key:nil, usage_tracker:nil, logger:nil,
|
|
50
|
+
http_connect_timeout:nil, http_read_timeout:nil, http_max_retries:nil)
|
|
51
|
+
|
|
52
|
+
@model = model || default_model
|
|
53
|
+
@api_key = api_key || default_api_key
|
|
54
|
+
@usage_tracker = usage_tracker
|
|
55
|
+
@logger = logger
|
|
56
|
+
@http_connect_timeout = http_connect_timeout
|
|
57
|
+
@http_read_timeout = http_read_timeout
|
|
58
|
+
@http_max_retries = http_max_retries
|
|
59
|
+
end
|
|
60
|
+
|
|
61
|
+
def name
|
|
62
|
+
raise NotImplementedError
|
|
63
|
+
end
|
|
64
|
+
|
|
65
|
+
def default_model
|
|
66
|
+
raise NotImplementedError
|
|
67
|
+
end
|
|
68
|
+
|
|
69
|
+
def default_api_key
|
|
70
|
+
raise NotImplementedError
|
|
71
|
+
end
|
|
72
|
+
|
|
73
|
+
def base_url
|
|
74
|
+
raise NotImplementedError
|
|
75
|
+
end
|
|
76
|
+
|
|
77
|
+
def parse_usage(response)
|
|
78
|
+
raise NotImplementedError
|
|
79
|
+
end
|
|
80
|
+
|
|
81
|
+
def build_client
|
|
82
|
+
client_class.new(
|
|
83
|
+
provider: self,
|
|
84
|
+
api_key: api_key,
|
|
85
|
+
logger: logger,
|
|
86
|
+
usage_tracker: usage_tracker,
|
|
87
|
+
http_connect_timeout: http_connect_timeout,
|
|
88
|
+
http_read_timeout: http_read_timeout,
|
|
89
|
+
http_max_retries: http_max_retries
|
|
90
|
+
)
|
|
91
|
+
end
|
|
92
|
+
|
|
93
|
+
def create_assistant(state:nil, system:nil, model:nil, tools:[], mcp_servers:{},
|
|
94
|
+
json_schema:nil, thinking:nil, client:nil)
|
|
95
|
+
|
|
96
|
+
raise ArgumentError, 'Use state or system, not both' if state && system
|
|
97
|
+
|
|
98
|
+
state = AssistantState.new(context: system) if state.nil? && system
|
|
99
|
+
|
|
100
|
+
assistant_class.new(
|
|
101
|
+
provider: self,
|
|
102
|
+
client: client,
|
|
103
|
+
model: model || self.model,
|
|
104
|
+
state: state,
|
|
105
|
+
tools: tools,
|
|
106
|
+
mcp_servers: mcp_servers,
|
|
107
|
+
json_schema: json_schema,
|
|
108
|
+
thinking: thinking,
|
|
109
|
+
logger: logger
|
|
110
|
+
)
|
|
111
|
+
end
|
|
112
|
+
|
|
113
|
+
def generate_text(prompt:nil, messages:nil, system:nil, model:nil,
|
|
114
|
+
json_schema:nil, thinking:nil, client:nil)
|
|
115
|
+
|
|
116
|
+
raise ArgumentError, 'Use prompt or messages, not both' if prompt && messages
|
|
117
|
+
raise ArgumentError, 'Undefined prompt or messages' if prompt.nil? && messages.nil?
|
|
118
|
+
|
|
119
|
+
system_prompt, conversation = split_system(messages || [{role: Roles::USER, content: prompt}])
|
|
120
|
+
|
|
121
|
+
client ||= build_client
|
|
122
|
+
|
|
123
|
+
response = request(
|
|
124
|
+
client: client,
|
|
125
|
+
messages: conversation.map { |message| encode_message message },
|
|
126
|
+
system: system || system_prompt,
|
|
127
|
+
model: model || self.model,
|
|
128
|
+
json_schema: json_schema,
|
|
129
|
+
thinking: thinking_config(thinking)
|
|
130
|
+
)
|
|
131
|
+
|
|
132
|
+
Result.new content: parse_content(response),
|
|
133
|
+
usage: parse_usage(response),
|
|
134
|
+
raw: response
|
|
135
|
+
end
|
|
136
|
+
|
|
137
|
+
private
|
|
138
|
+
|
|
139
|
+
attr_reader :api_key, :usage_tracker, :logger,
|
|
140
|
+
:http_connect_timeout, :http_read_timeout, :http_max_retries
|
|
141
|
+
|
|
142
|
+
def provider_module
|
|
143
|
+
AI.const_get MODULES.fetch(name)
|
|
144
|
+
end
|
|
145
|
+
|
|
146
|
+
def client_class
|
|
147
|
+
provider_module::Client
|
|
148
|
+
end
|
|
149
|
+
|
|
150
|
+
def assistant_class
|
|
151
|
+
provider_module::Assistant
|
|
152
|
+
end
|
|
153
|
+
|
|
154
|
+
def split_system(messages)
|
|
155
|
+
normalized = messages.map { |message| normalize_message message }
|
|
156
|
+
|
|
157
|
+
system_messages, conversation = normalized.partition { |message| message[:role] == Roles::SYSTEM }
|
|
158
|
+
|
|
159
|
+
system_prompt = system_messages.map { |message| message[:content] }.join("\n") unless system_messages.empty?
|
|
160
|
+
|
|
161
|
+
[system_prompt, conversation]
|
|
162
|
+
end
|
|
163
|
+
|
|
164
|
+
def normalize_message(message)
|
|
165
|
+
{
|
|
166
|
+
role: (message[:role] || message['role'] || Roles::USER).to_s,
|
|
167
|
+
content: message[:content] || message['content']
|
|
168
|
+
}
|
|
169
|
+
end
|
|
170
|
+
|
|
171
|
+
def thinking_config(level)
|
|
172
|
+
return nil if level.nil?
|
|
173
|
+
|
|
174
|
+
levels = self.class::THINKING_LEVELS
|
|
175
|
+
|
|
176
|
+
raise ArgumentError, "Invalid thinking level '#{level}'. Valid: #{levels.keys.join(', ')}" unless levels.key? level
|
|
177
|
+
|
|
178
|
+
levels[level]
|
|
179
|
+
end
|
|
180
|
+
|
|
181
|
+
# --- Template methods ---
|
|
182
|
+
|
|
183
|
+
def request(client:, messages:, system:, model:, json_schema:, thinking:)
|
|
184
|
+
raise NotImplementedError
|
|
185
|
+
end
|
|
186
|
+
|
|
187
|
+
def encode_message(message)
|
|
188
|
+
raise NotImplementedError
|
|
189
|
+
end
|
|
190
|
+
|
|
191
|
+
def parse_content(response)
|
|
192
|
+
raise NotImplementedError
|
|
193
|
+
end
|
|
194
|
+
|
|
195
|
+
end
|
|
196
|
+
end
|
|
197
|
+
end
|
|
@@ -0,0 +1,17 @@
|
|
|
1
|
+
module Rasti
|
|
2
|
+
module AI
|
|
3
|
+
module ProviderAware
|
|
4
|
+
|
|
5
|
+
private
|
|
6
|
+
|
|
7
|
+
def provider
|
|
8
|
+
@provider ||= provider_module::Provider.new
|
|
9
|
+
end
|
|
10
|
+
|
|
11
|
+
def provider_module
|
|
12
|
+
Object.const_get self.class.name.split('::')[0..-2].join('::')
|
|
13
|
+
end
|
|
14
|
+
|
|
15
|
+
end
|
|
16
|
+
end
|
|
17
|
+
end
|
data/lib/rasti/ai/version.rb
CHANGED
data/lib/rasti/ai.rb
CHANGED
|
@@ -15,12 +15,17 @@ module Rasti
|
|
|
15
15
|
extend ClassConfig
|
|
16
16
|
|
|
17
17
|
require_relative 'ai/errors'
|
|
18
|
+
require_relative 'ai/roles'
|
|
18
19
|
require_relative 'ai/usage'
|
|
20
|
+
require_relative 'ai/result'
|
|
19
21
|
require_relative 'ai/assistant_state'
|
|
20
22
|
require_relative 'ai/tool'
|
|
21
23
|
require_relative 'ai/tool_serializer'
|
|
24
|
+
require_relative 'ai/provider_aware'
|
|
22
25
|
require_relative 'ai/client'
|
|
23
26
|
require_relative 'ai/assistant'
|
|
27
|
+
require_relative 'ai/provider'
|
|
28
|
+
require_relative_pattern 'ai/open_ai/*'
|
|
24
29
|
require_relative_pattern 'ai/**/*'
|
|
25
30
|
|
|
26
31
|
attr_config :logger, Logger.new(STDOUT)
|
|
@@ -29,6 +34,8 @@ module Rasti
|
|
|
29
34
|
attr_config :http_read_timeout, 60
|
|
30
35
|
attr_config :http_max_retries, 3
|
|
31
36
|
|
|
37
|
+
attr_config :default_provider, ENV['AI_DEFAULT_PROVIDER']
|
|
38
|
+
|
|
32
39
|
attr_config :openai_api_key, ENV['OPENAI_API_KEY']
|
|
33
40
|
attr_config :openai_default_model, ENV['OPENAI_DEFAULT_MODEL']
|
|
34
41
|
|
|
@@ -38,7 +45,100 @@ module Rasti
|
|
|
38
45
|
attr_config :anthropic_api_key, ENV['ANTHROPIC_API_KEY']
|
|
39
46
|
attr_config :anthropic_default_model, ENV['ANTHROPIC_DEFAULT_MODEL']
|
|
40
47
|
|
|
48
|
+
attr_config :openrouter_api_key, ENV['OPENROUTER_API_KEY']
|
|
49
|
+
attr_config :openrouter_default_model, ENV['OPENROUTER_DEFAULT_MODEL']
|
|
50
|
+
|
|
51
|
+
attr_config :huawei_maas_api_key, ENV['HUAWEI_MAAS_API_KEY']
|
|
52
|
+
attr_config :huawei_maas_default_model, ENV['HUAWEI_MAAS_DEFAULT_MODEL']
|
|
53
|
+
|
|
41
54
|
attr_config :usage_tracker, nil
|
|
42
55
|
|
|
56
|
+
class << self
|
|
57
|
+
|
|
58
|
+
def provider(name=nil, model:nil, api_key:nil, usage_tracker:nil, logger:nil,
|
|
59
|
+
http_connect_timeout:nil, http_read_timeout:nil, http_max_retries:nil)
|
|
60
|
+
|
|
61
|
+
build_provider(
|
|
62
|
+
name, model, api_key, usage_tracker, logger,
|
|
63
|
+
http_connect_timeout, http_read_timeout, http_max_retries
|
|
64
|
+
).first
|
|
65
|
+
end
|
|
66
|
+
|
|
67
|
+
def create_assistant(provider:nil, model:nil, api_key:nil, usage_tracker:nil, logger:nil,
|
|
68
|
+
http_connect_timeout:nil, http_read_timeout:nil, http_max_retries:nil,
|
|
69
|
+
state:nil, system:nil, tools:[], mcp_servers:{}, json_schema:nil,
|
|
70
|
+
thinking:nil, client:nil)
|
|
71
|
+
|
|
72
|
+
resolved_provider, resolved_model = build_provider(
|
|
73
|
+
provider, model, api_key, usage_tracker, logger,
|
|
74
|
+
http_connect_timeout, http_read_timeout, http_max_retries
|
|
75
|
+
)
|
|
76
|
+
|
|
77
|
+
resolved_provider.create_assistant(
|
|
78
|
+
model: resolved_model,
|
|
79
|
+
state: state,
|
|
80
|
+
system: system,
|
|
81
|
+
tools: tools,
|
|
82
|
+
mcp_servers: mcp_servers,
|
|
83
|
+
json_schema: json_schema,
|
|
84
|
+
thinking: thinking,
|
|
85
|
+
client: client
|
|
86
|
+
)
|
|
87
|
+
end
|
|
88
|
+
|
|
89
|
+
def generate_text(provider:nil, model:nil, api_key:nil, usage_tracker:nil, logger:nil,
|
|
90
|
+
http_connect_timeout:nil, http_read_timeout:nil, http_max_retries:nil,
|
|
91
|
+
prompt:nil, messages:nil, system:nil, json_schema:nil,
|
|
92
|
+
thinking:nil, client:nil)
|
|
93
|
+
|
|
94
|
+
resolved_provider, resolved_model = build_provider(
|
|
95
|
+
provider, model, api_key, usage_tracker, logger,
|
|
96
|
+
http_connect_timeout, http_read_timeout, http_max_retries
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
resolved_provider.generate_text(
|
|
100
|
+
model: resolved_model,
|
|
101
|
+
prompt: prompt,
|
|
102
|
+
messages: messages,
|
|
103
|
+
system: system,
|
|
104
|
+
json_schema: json_schema,
|
|
105
|
+
thinking: thinking,
|
|
106
|
+
client: client
|
|
107
|
+
)
|
|
108
|
+
end
|
|
109
|
+
|
|
110
|
+
private
|
|
111
|
+
|
|
112
|
+
def build_provider(name, model, api_key, usage_tracker, logger,
|
|
113
|
+
http_connect_timeout, http_read_timeout, http_max_retries)
|
|
114
|
+
|
|
115
|
+
name, model = split_provider_and_model name, model
|
|
116
|
+
|
|
117
|
+
resolved_provider = Provider.build name,
|
|
118
|
+
model: model,
|
|
119
|
+
api_key: api_key,
|
|
120
|
+
usage_tracker: usage_tracker,
|
|
121
|
+
logger: logger,
|
|
122
|
+
http_connect_timeout: http_connect_timeout,
|
|
123
|
+
http_read_timeout: http_read_timeout,
|
|
124
|
+
http_max_retries: http_max_retries
|
|
125
|
+
|
|
126
|
+
[resolved_provider, model]
|
|
127
|
+
end
|
|
128
|
+
|
|
129
|
+
def split_provider_and_model(name, model)
|
|
130
|
+
if name.nil? && model.to_s.include?(':')
|
|
131
|
+
name, model = model.split ':', 2
|
|
132
|
+
end
|
|
133
|
+
|
|
134
|
+
name ||= default_provider
|
|
135
|
+
|
|
136
|
+
raise ArgumentError, "Undefined provider. Set it with provider: or Rasti::AI.default_provider. Valid providers: #{Provider.names.join(', ')}" if name.nil?
|
|
137
|
+
|
|
138
|
+
[name, model]
|
|
139
|
+
end
|
|
140
|
+
|
|
141
|
+
end
|
|
142
|
+
|
|
43
143
|
end
|
|
44
|
-
end
|
|
144
|
+
end
|
|
@@ -0,0 +1,106 @@
|
|
|
1
|
+
require 'minitest_helper'
|
|
2
|
+
|
|
3
|
+
describe Rasti::AI::Anthropic::Provider do
|
|
4
|
+
|
|
5
|
+
let(:provider) { Rasti::AI.provider :anthropic }
|
|
6
|
+
|
|
7
|
+
let(:question) { 'How many goals has Messi scored for Barca?' }
|
|
8
|
+
|
|
9
|
+
let(:answer) { 'Lionel Messi scored 672 goals in 778 official matches for FC Barcelona.' }
|
|
10
|
+
|
|
11
|
+
def stub_messages(request_body, response_body=nil)
|
|
12
|
+
stub_request(:post, 'https://api.anthropic.com/v1/messages')
|
|
13
|
+
.with(headers: {'x-api-key' => Rasti::AI.anthropic_api_key}, body: JSON.dump(request_body))
|
|
14
|
+
.to_return(body: response_body || read_resource('anthropic/basic_response.json', content: answer))
|
|
15
|
+
end
|
|
16
|
+
|
|
17
|
+
it 'Name, model and classes' do
|
|
18
|
+
assert_equal :anthropic, provider.name
|
|
19
|
+
assert_equal Rasti::AI.anthropic_default_model, provider.model
|
|
20
|
+
assert_instance_of Rasti::AI::Anthropic::Client, provider.build_client
|
|
21
|
+
assert_instance_of Rasti::AI::Anthropic::Assistant, provider.create_assistant
|
|
22
|
+
end
|
|
23
|
+
|
|
24
|
+
it 'Generate text' do
|
|
25
|
+
stub_messages model: Rasti::AI.anthropic_default_model,
|
|
26
|
+
max_tokens: 4096,
|
|
27
|
+
messages: [
|
|
28
|
+
{
|
|
29
|
+
role: 'user',
|
|
30
|
+
content: question
|
|
31
|
+
}
|
|
32
|
+
]
|
|
33
|
+
|
|
34
|
+
result = provider.generate_text prompt: question
|
|
35
|
+
|
|
36
|
+
assert_equal answer, result.content
|
|
37
|
+
assert_equal 'anthropic', result.usage.provider
|
|
38
|
+
assert_equal 'end_turn', result.raw['stop_reason']
|
|
39
|
+
end
|
|
40
|
+
|
|
41
|
+
it 'Generate text with system, thinking and model override' do
|
|
42
|
+
stub_messages model: 'claude-sonnet-4-5',
|
|
43
|
+
max_tokens: 4096,
|
|
44
|
+
messages: [
|
|
45
|
+
{
|
|
46
|
+
role: 'user',
|
|
47
|
+
content: question
|
|
48
|
+
}
|
|
49
|
+
],
|
|
50
|
+
thinking: {type: 'enabled', budget_tokens: 16_000},
|
|
51
|
+
system: 'Act as sports journalist'
|
|
52
|
+
|
|
53
|
+
result = provider.generate_text prompt: question,
|
|
54
|
+
model: 'claude-sonnet-4-5',
|
|
55
|
+
system: 'Act as sports journalist',
|
|
56
|
+
thinking: 'high'
|
|
57
|
+
|
|
58
|
+
assert_equal answer, result.content
|
|
59
|
+
end
|
|
60
|
+
|
|
61
|
+
it 'Generate text with json schema' do
|
|
62
|
+
json_schema = {player: 'string'}
|
|
63
|
+
|
|
64
|
+
structured_response = {
|
|
65
|
+
'content' => [
|
|
66
|
+
{
|
|
67
|
+
'type' => 'tool_use',
|
|
68
|
+
'id' => 'toolu_1',
|
|
69
|
+
'name' => 'structured_output',
|
|
70
|
+
'input' => {'player' => 'Lionel Messi'}
|
|
71
|
+
}
|
|
72
|
+
],
|
|
73
|
+
'stop_reason' => 'tool_use'
|
|
74
|
+
}
|
|
75
|
+
|
|
76
|
+
stub_messages(
|
|
77
|
+
{
|
|
78
|
+
model: Rasti::AI.anthropic_default_model,
|
|
79
|
+
max_tokens: 4096,
|
|
80
|
+
messages: [
|
|
81
|
+
{
|
|
82
|
+
role: 'user',
|
|
83
|
+
content: question
|
|
84
|
+
}
|
|
85
|
+
],
|
|
86
|
+
tools: [
|
|
87
|
+
{
|
|
88
|
+
name: 'structured_output',
|
|
89
|
+
description: 'Return the structured response',
|
|
90
|
+
input_schema: {
|
|
91
|
+
type: 'object',
|
|
92
|
+
properties: json_schema
|
|
93
|
+
}
|
|
94
|
+
}
|
|
95
|
+
],
|
|
96
|
+
tool_choice: {type: 'tool', name: 'structured_output'}
|
|
97
|
+
},
|
|
98
|
+
JSON.dump(structured_response)
|
|
99
|
+
)
|
|
100
|
+
|
|
101
|
+
result = provider.generate_text prompt: question, json_schema: json_schema
|
|
102
|
+
|
|
103
|
+
assert_equal '{"player":"Lionel Messi"}', result.content
|
|
104
|
+
end
|
|
105
|
+
|
|
106
|
+
end
|
|
@@ -0,0 +1,90 @@
|
|
|
1
|
+
require 'minitest_helper'
|
|
2
|
+
|
|
3
|
+
describe Rasti::AI::Gemini::Provider do
|
|
4
|
+
|
|
5
|
+
let(:provider) { Rasti::AI.provider :gemini }
|
|
6
|
+
|
|
7
|
+
let(:question) { 'How many goals has Messi scored for Barca?' }
|
|
8
|
+
|
|
9
|
+
let(:answer) { 'Lionel Messi scored 672 goals in 778 official matches for FC Barcelona.' }
|
|
10
|
+
|
|
11
|
+
def api_url(model:nil)
|
|
12
|
+
"https://generativelanguage.googleapis.com/v1beta/models/#{model || Rasti::AI.gemini_default_model}:generateContent?key=#{Rasti::AI.gemini_api_key}"
|
|
13
|
+
end
|
|
14
|
+
|
|
15
|
+
def stub_generate_content(request_body)
|
|
16
|
+
stub_request(:post, api_url)
|
|
17
|
+
.with(body: JSON.dump(request_body))
|
|
18
|
+
.to_return(body: read_resource('gemini/basic_response.json', content: answer))
|
|
19
|
+
end
|
|
20
|
+
|
|
21
|
+
it 'Name, model and classes' do
|
|
22
|
+
assert_equal :gemini, provider.name
|
|
23
|
+
assert_equal Rasti::AI.gemini_default_model, provider.model
|
|
24
|
+
assert_instance_of Rasti::AI::Gemini::Client, provider.build_client
|
|
25
|
+
assert_instance_of Rasti::AI::Gemini::Assistant, provider.create_assistant
|
|
26
|
+
end
|
|
27
|
+
|
|
28
|
+
it 'Generate text' do
|
|
29
|
+
stub_generate_content contents: [
|
|
30
|
+
{
|
|
31
|
+
role: 'user',
|
|
32
|
+
parts: [{text: question}]
|
|
33
|
+
}
|
|
34
|
+
]
|
|
35
|
+
|
|
36
|
+
result = provider.generate_text prompt: question
|
|
37
|
+
|
|
38
|
+
assert_equal answer, result.content
|
|
39
|
+
assert_equal 'gemini', result.usage.provider
|
|
40
|
+
assert_equal 'gemini-test', result.raw['modelVersion']
|
|
41
|
+
end
|
|
42
|
+
|
|
43
|
+
it 'Generate text with generic roles and system' do
|
|
44
|
+
stub_generate_content contents: [
|
|
45
|
+
{
|
|
46
|
+
role: 'user',
|
|
47
|
+
parts: [{text: 'who is the best player'}]
|
|
48
|
+
},
|
|
49
|
+
{
|
|
50
|
+
role: 'model',
|
|
51
|
+
parts: [{text: 'Lionel Messi'}]
|
|
52
|
+
},
|
|
53
|
+
{
|
|
54
|
+
role: 'user',
|
|
55
|
+
parts: [{text: question}]
|
|
56
|
+
}
|
|
57
|
+
],
|
|
58
|
+
system_instruction: {
|
|
59
|
+
parts: [{text: 'Act as sports journalist'}]
|
|
60
|
+
}
|
|
61
|
+
|
|
62
|
+
messages = [
|
|
63
|
+
{role: Rasti::AI::Roles::SYSTEM, content: 'Act as sports journalist'},
|
|
64
|
+
{role: Rasti::AI::Roles::USER, content: 'who is the best player'},
|
|
65
|
+
{role: Rasti::AI::Roles::ASSISTANT, content: 'Lionel Messi'},
|
|
66
|
+
{role: Rasti::AI::Roles::USER, content: question}
|
|
67
|
+
]
|
|
68
|
+
|
|
69
|
+
assert_equal answer, provider.generate_text(messages: messages).content
|
|
70
|
+
end
|
|
71
|
+
|
|
72
|
+
it 'Generate text with json schema and thinking' do
|
|
73
|
+
json_schema = {type: 'object', properties: {player: {type: 'string'}}}
|
|
74
|
+
|
|
75
|
+
stub_generate_content contents: [
|
|
76
|
+
{
|
|
77
|
+
role: 'user',
|
|
78
|
+
parts: [{text: question}]
|
|
79
|
+
}
|
|
80
|
+
],
|
|
81
|
+
generation_config: {
|
|
82
|
+
thinking_config: {thinking_budget: 8_192},
|
|
83
|
+
response_mime_type: 'application/json',
|
|
84
|
+
response_schema: json_schema
|
|
85
|
+
}
|
|
86
|
+
|
|
87
|
+
assert_equal answer, provider.generate_text(prompt: question, json_schema: json_schema, thinking: 'medium').content
|
|
88
|
+
end
|
|
89
|
+
|
|
90
|
+
end
|
|
@@ -0,0 +1,80 @@
|
|
|
1
|
+
require 'minitest_helper'
|
|
2
|
+
|
|
3
|
+
describe Rasti::AI::HuaweiMaaS::Assistant do
|
|
4
|
+
|
|
5
|
+
let(:api_url) { 'https://api-ap-southeast-1.modelarts-maas.com/v2/chat/completions' }
|
|
6
|
+
|
|
7
|
+
let(:question) { 'How many goals has Messi scored for Barca?' }
|
|
8
|
+
|
|
9
|
+
let(:answer) { 'Lionel Messi scored 672 goals in 778 official matches for FC Barcelona.' }
|
|
10
|
+
|
|
11
|
+
it 'Default' do
|
|
12
|
+
stub_request(:post, api_url)
|
|
13
|
+
.with(body: read_resource('open_ai/basic_request.json', model: Rasti::AI.huawei_maas_default_model, prompt: question))
|
|
14
|
+
.to_return(body: read_resource('open_ai/basic_response.json', content: answer))
|
|
15
|
+
|
|
16
|
+
assistant = Rasti::AI::HuaweiMaaS::Assistant.new
|
|
17
|
+
|
|
18
|
+
response = assistant.call question
|
|
19
|
+
|
|
20
|
+
assert_equal answer, response
|
|
21
|
+
end
|
|
22
|
+
|
|
23
|
+
describe 'Tools' do
|
|
24
|
+
|
|
25
|
+
let(:client) { Minitest::Mock.new }
|
|
26
|
+
|
|
27
|
+
let(:tool_response) do
|
|
28
|
+
read_json_resource(
|
|
29
|
+
'open_ai/tool_response.json',
|
|
30
|
+
name: 'goals_by_player',
|
|
31
|
+
arguments: {player: 'Lionel Messi', team: 'Barcelona'}
|
|
32
|
+
)
|
|
33
|
+
end
|
|
34
|
+
|
|
35
|
+
let(:tool_result) { '672' }
|
|
36
|
+
|
|
37
|
+
let(:answer_with_tool) { 'Lionel Messi scored 672 goals for FC Barcelona.' }
|
|
38
|
+
|
|
39
|
+
def basic_response(content)
|
|
40
|
+
read_json_resource('open_ai/basic_response.json', content: content)
|
|
41
|
+
end
|
|
42
|
+
|
|
43
|
+
def stub_client_request(role:, content:, response:, tools:[])
|
|
44
|
+
serialized_tools = tools.map do |tool|
|
|
45
|
+
{type: 'function', function: Rasti::AI::ToolSerializer.serialize(tool.class)}
|
|
46
|
+
end
|
|
47
|
+
|
|
48
|
+
client.expect :chat_completions, response do |params|
|
|
49
|
+
last_message = params[:messages].last
|
|
50
|
+
last_message[:role] == role &&
|
|
51
|
+
last_message[:content] == content &&
|
|
52
|
+
params[:tools] == serialized_tools
|
|
53
|
+
end
|
|
54
|
+
end
|
|
55
|
+
|
|
56
|
+
it 'Call tool' do
|
|
57
|
+
tool = GoalsByPlayer.new
|
|
58
|
+
|
|
59
|
+
stub_client_request role: Rasti::AI::HuaweiMaaS::Roles::USER,
|
|
60
|
+
content: question,
|
|
61
|
+
tools: [tool],
|
|
62
|
+
response: tool_response
|
|
63
|
+
|
|
64
|
+
stub_client_request role: Rasti::AI::HuaweiMaaS::Roles::TOOL,
|
|
65
|
+
content: tool_result,
|
|
66
|
+
tools: [tool],
|
|
67
|
+
response: basic_response(answer_with_tool)
|
|
68
|
+
|
|
69
|
+
assistant = Rasti::AI::HuaweiMaaS::Assistant.new client: client, tools: [tool]
|
|
70
|
+
|
|
71
|
+
response = assistant.call question
|
|
72
|
+
|
|
73
|
+
assert_equal answer_with_tool, response
|
|
74
|
+
|
|
75
|
+
client.verify
|
|
76
|
+
end
|
|
77
|
+
|
|
78
|
+
end
|
|
79
|
+
|
|
80
|
+
end
|