ruby_decision_model 0.0.1 → 0.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.
@@ -2,137 +2,242 @@
2
2
 
3
3
  require "json"
4
4
  require "net/http"
5
+ require "openssl"
5
6
  require "uri"
6
7
 
7
8
  module RubyDecisionModel
8
9
  class Client
9
- DEFAULT_BASE_URL = "https://openrouter.ai/api/alpha"
10
- DEFAULT_MODEL = "typesafe/jev-1.13"
11
- MAX_ATTEMPTS = 2
12
- RETRYABLE_STATUSES = [429, 500, 502, 503, 504, 524, 529].freeze
13
- RETRYABLE_EXCEPTIONS = [
14
- Net::OpenTimeout,
15
- Net::ReadTimeout,
16
- Errno::ECONNRESET,
17
- Errno::ECONNREFUSED,
18
- Errno::EPIPE,
19
- SocketError,
20
- IOError
21
- ].freeze
22
-
23
- def initialize(api_key:, model: DEFAULT_MODEL, base_url: DEFAULT_BASE_URL, timeout: 5,
24
- transport: nil, sleeper: ->(seconds) { sleep(seconds) })
25
- raise ConfigurationError, "api_key is required" if api_key.nil? || api_key.to_s.strip.empty?
26
- raise ConfigurationError, "model is required" if model.nil? || model.to_s.strip.empty?
27
- raise ConfigurationError, "base_url is required" if base_url.nil? || base_url.to_s.strip.empty?
28
-
29
- @api_key = api_key
30
- @model = model
31
- @base_url = base_url.to_s.chomp("/")
32
- @timeout = timeout
10
+ # Kept from 0.0.1 for callers that referenced them. They describe the
11
+ # Typesafe and OpenRouter providers and the default RetryPolicy; prefer
12
+ # those directly. Each provider now names its own request id header.
13
+ REQUEST_ID_HEADER = "x-typesafe-request-id"
14
+ DEFAULT_BASE_URL = Providers::OpenRouter.new.default_base_url
15
+ DEFAULT_MODEL = Providers::OpenRouter.new.default_model
16
+ MAX_ATTEMPTS = RetryPolicy.new.max_retries + 1
17
+ RETRYABLE_STATUSES = RetryPolicy::DEFAULT_STATUSES
18
+ RETRYABLE_EXCEPTIONS = (RetryPolicy::TIMEOUT_EXCEPTIONS + RetryPolicy::CONNECTION_EXCEPTIONS).freeze
19
+
20
+ # Connecting is quick wherever the answer is slow, so the open timeout
21
+ # stays short unless the caller sets timeout: explicitly.
22
+ DEFAULT_OPEN_TIMEOUT = 5
23
+
24
+ attr_reader :provider, :model, :retry_policy, :timeout, :open_timeout
25
+
26
+ # provider: a name from Providers.names or a Providers::Base instance.
27
+ # When nil, RUBY_DECISION_MODEL_PROVIDER names one if set;
28
+ # otherwise api_key: alone selects OpenRouter, and with no
29
+ # api_key: the environment decides (see Providers.from_env).
30
+ # api_key: overrides the provider's env var.
31
+ # model: nil means the provider default; aliases resolve per provider.
32
+ # base_url: overrides the provider base URL.
33
+ # timeout: read timeout in seconds, and open timeout too when given;
34
+ # nil means the provider's read default (5 for Jev-speed APIs,
35
+ # 30 where the vendor documents multi-second responses) with a
36
+ # 5 second open timeout.
37
+ # transport: callable(url:, headers:, body:) returning
38
+ # [status, body_string, headers_hash] (a 2-element return is
39
+ # still accepted and treated as having no headers).
40
+ # retry: a RetryPolicy or a Hash of overrides.
41
+ # random: callable returning a Float in 0...1, used for backoff jitter.
42
+ # clock: callable returning monotonic seconds, used for total_timeout.
43
+ def initialize(provider: nil, api_key: nil, model: nil, base_url: nil, timeout: nil,
44
+ transport: nil, sleeper: ->(seconds) { sleep(seconds) }, retry: {},
45
+ random: -> { rand }, clock: -> { Process.clock_gettime(Process::CLOCK_MONOTONIC) })
46
+ @provider = resolve_provider(provider, api_key: api_key, base_url: base_url)
47
+ @provider.validate!
48
+
49
+ @model = @provider.resolve_model(model)
50
+ @timeout = timeout || @provider.default_timeout
51
+ @open_timeout = timeout || DEFAULT_OPEN_TIMEOUT
33
52
  @transport = transport || default_transport
34
53
  @sleeper = sleeper
54
+ @retry_policy = RetryPolicy.from(binding.local_variable_get(:retry))
55
+ @random = random
56
+ @clock = clock
35
57
  end
36
58
 
37
- def ask(state:, questions:)
38
- raise RequestError, "questions must not be empty" if questions.nil? || questions.empty?
59
+ def base_url
60
+ @provider.base_url
61
+ end
39
62
 
40
- body = JSON.generate({ "model" => @model, "state" => state, "questions" => questions })
41
- headers = {
42
- "Authorization" => "Bearer #{@api_key}",
43
- "Content-Type" => "application/json",
44
- "Accept" => "application/json"
45
- }
63
+ # images: data URLs (see Images) for providers that read images. Each
64
+ # provider places them where its API expects.
65
+ def ask(state:, questions:, images: nil)
66
+ raise RequestError, "questions must not be empty" if questions.nil? || questions.empty?
46
67
 
47
- status, response_body = perform_with_retry(url: "#{@base_url}/decisions", headers: headers, body: body)
48
- handle_response(status, response_body, questions)
68
+ body = build_body(state, questions, images)
69
+ status, response_body, response_headers = perform_with_retry(
70
+ url: request_url, headers: @provider.headers, body: body
71
+ )
72
+ handle_response(status, response_body, response_headers, questions)
49
73
  end
50
74
 
51
75
  private
52
76
 
77
+ # Providers written against 0.1.0 may override url without the model
78
+ # argument added in 0.2.0.
79
+ def request_url
80
+ takes_model = @provider.method(:url).parameters.any? { |type, _| %i[req opt rest].include?(type) }
81
+ takes_model ? @provider.url(@model) : @provider.url
82
+ end
83
+
84
+ def build_body(state, questions, images)
85
+ images = Array(images)
86
+ return @provider.request_body(model: @model, state: state, questions: questions) if images.empty?
87
+
88
+ raise RequestError, "#{@provider.name} does not accept images" unless @provider.supports_images?
89
+
90
+ images.each do |image|
91
+ next if image.is_a?(String) && image.start_with?("data:image/")
92
+
93
+ raise RequestError, "images must be data URLs (data:image/...); see RubyDecisionModel::Images"
94
+ end
95
+
96
+ @provider.request_body(model: @model, state: state, questions: questions, images: images)
97
+ rescue JSON::GeneratorError, JSON::NestingError => e
98
+ # The generator's message can quote the offending value, so it stays
99
+ # on #cause rather than in a message that may be logged.
100
+ raise RequestError, "state or questions could not be encoded as JSON (#{e.class})"
101
+ end
102
+
103
+ def resolve_provider(provider, api_key:, base_url:)
104
+ case provider
105
+ when Providers::Base
106
+ return provider if api_key.nil? && base_url.nil?
107
+
108
+ # Never mutate a provider the caller may share between clients.
109
+ provider.dup.configure(api_key: api_key, base_url: base_url)
110
+ when Symbol, String
111
+ Providers.build(provider, api_key: api_key, base_url: base_url)
112
+ when nil
113
+ if api_key.nil?
114
+ Providers.from_env&.configure(base_url: base_url) || raise(
115
+ ConfigurationError,
116
+ "no provider configured: pass provider: or api_key:, or set one of #{Providers.env_vars.join(', ')}"
117
+ )
118
+ else
119
+ Providers.build(Providers.named_in_env || :open_router, api_key: api_key, base_url: base_url)
120
+ end
121
+ else
122
+ raise ConfigurationError, "provider must be a Symbol or a Providers::Base, got #{provider.class}"
123
+ end
124
+ end
125
+
53
126
  def default_transport
54
127
  lambda do |url:, headers:, body:|
55
128
  uri = URI.parse(url)
56
129
  http = Net::HTTP.new(uri.host, uri.port)
57
130
  http.use_ssl = uri.scheme == "https"
58
- http.open_timeout = @timeout
131
+ http.open_timeout = @open_timeout
59
132
  http.read_timeout = @timeout
133
+ http.write_timeout = @timeout
60
134
 
61
135
  request = Net::HTTP::Post.new(uri.request_uri)
62
136
  headers.each { |k, v| request[k] = v }
63
137
  request.body = body
64
138
 
65
139
  response = http.request(request)
66
- [response.code.to_i, response.body]
140
+ [response.code.to_i, response.body, response.each_header.to_h]
67
141
  end
68
142
  end
69
143
 
70
144
  def perform_with_retry(url:, headers:, body:)
71
- attempts = 0
145
+ policy = @retry_policy
146
+ started_at = @clock.call
147
+ retries = 0
72
148
 
73
149
  loop do
74
- attempts += 1
75
150
  begin
76
- status, response_body = @transport.call(url: url, headers: headers, body: body)
77
- rescue *RETRYABLE_EXCEPTIONS => e
78
- raise_transport_error(e) if attempts >= MAX_ATTEMPTS
79
-
80
- @sleeper.call(backoff_seconds)
81
- next
151
+ status, response_body, response_headers = normalize_transport_result(
152
+ @transport.call(url: url, headers: headers, body: body)
153
+ )
82
154
  rescue Error
83
155
  raise
84
156
  rescue StandardError => e
85
- raise_transport_error(e)
157
+ raise_transport_error(e) unless policy.retryable_exception?(e) && retries < policy.max_retries
158
+
159
+ delay = policy.backoff(retries, random: @random)
160
+ raise_transport_error(e) if budget_exceeded?(policy, started_at, delay)
161
+
162
+ @sleeper.call(delay)
163
+ raise_transport_error(e) if budget_exceeded?(policy, started_at, 0.0)
164
+
165
+ retries += 1
166
+ next
86
167
  end
87
168
 
88
- return [status, response_body] unless RETRYABLE_STATUSES.include?(status) && attempts < MAX_ATTEMPTS
169
+ result = [status, response_body, response_headers]
170
+ return result unless policy.retryable_status?(status) && retries < policy.max_retries
171
+
172
+ delay = policy.delay(retries, headers: response_headers, random: @random)
173
+ return result if budget_exceeded?(policy, started_at, delay)
174
+
175
+ @sleeper.call(delay)
176
+ return result if budget_exceeded?(policy, started_at, 0.0)
89
177
 
90
- @sleeper.call(backoff_seconds)
178
+ retries += 1
91
179
  end
92
180
  end
93
181
 
94
- def backoff_seconds
95
- 0.5 + (rand * 0.25)
182
+ def budget_exceeded?(policy, started_at, delay)
183
+ return false if policy.total_timeout.nil?
184
+
185
+ (@clock.call - started_at) + delay > policy.total_timeout
186
+ end
187
+
188
+ def normalize_transport_result(result)
189
+ status, response_body, response_headers = Array(result)
190
+ [status, response_body, response_headers.is_a?(Hash) ? response_headers : {}]
96
191
  end
97
192
 
98
193
  def raise_transport_error(exception)
99
- if exception.is_a?(Net::OpenTimeout) || exception.is_a?(Net::ReadTimeout)
194
+ if @retry_policy.timeout_exception?(exception)
100
195
  raise TimeoutError.new("request timed out: #{exception.message}", cause_error: exception)
101
196
  end
102
197
 
103
198
  raise TransportError.new("transport error: #{exception.message}", cause_error: exception)
104
199
  end
105
200
 
106
- def handle_response(status, response_body, questions)
107
- case status
108
- when 200..299
109
- parse_success(response_body, questions)
110
- when 401
111
- raise Unauthorized.new("unauthorized", status: status, body: response_body)
112
- when 413
113
- raise PayloadTooLarge.new("payload too large", status: status, body: response_body)
114
- when 429
115
- raise RateLimited.new("rate limited", status: status, body: response_body)
116
- else
117
- raise ApiError.new("api error (status #{status})", status: status, body: response_body)
118
- end
201
+ ERROR_CLASSES = {
202
+ 401 => [Unauthorized, "unauthorized"],
203
+ 413 => [PayloadTooLarge, "payload too large"],
204
+ 422 => [UnprocessableEntity, "unprocessable entity"],
205
+ 429 => [RateLimited, "rate limited"],
206
+ 529 => [Overloaded, "overloaded"]
207
+ }.freeze
208
+ private_constant :ERROR_CLASSES
209
+
210
+ def handle_response(status, response_body, response_headers, questions)
211
+ return parse_success(response_body, response_headers, questions) if (200..299).cover?(status)
212
+
213
+ klass, message = ERROR_CLASSES.fetch(status) { [ApiError, "api error (status #{status})"] }
214
+ detail = @provider.error_message(response_body)
215
+ message = "#{message}: #{detail}" if detail
216
+
217
+ raise klass.new(message, status: status, body: response_body, headers: response_headers)
119
218
  end
120
219
 
121
- def parse_success(response_body, questions)
220
+ def parse_success(response_body, response_headers, questions)
221
+ raise InvalidResponse, "response body was empty" if response_body.nil? || response_body.to_s.strip.empty?
222
+
122
223
  parsed = begin
123
- JSON.parse(response_body)
224
+ JSON.parse(response_body.to_s)
124
225
  rescue JSON::ParserError => e
125
226
  raise InvalidResponse, "response body was not valid JSON: #{e.message}"
126
227
  end
127
228
 
128
229
  raise InvalidResponse, "response body was not a JSON object" unless parsed.is_a?(Hash)
129
230
 
130
- raw_answers = parsed["answers"]
231
+ canonical = @provider.normalize_response(parsed, questions: questions)
232
+ raise InvalidResponse, "response body was not a JSON object" unless canonical.is_a?(Hash)
233
+
234
+ raw_answers = canonical["answers"]
131
235
  raw_answers = {} unless raw_answers.is_a?(Hash)
132
236
 
133
237
  normalized = {}
134
238
  malformed = []
135
239
  missing = []
240
+ refused = []
136
241
 
137
242
  questions.each do |raw_id, question|
138
243
  id = raw_id.to_s
@@ -146,6 +251,7 @@ module RubyDecisionModel
146
251
  malformed << id
147
252
  end
148
253
  else
254
+ refused << id if answer_hash.is_a?(Hash) && answer_hash["type"] == "refusal"
149
255
  missing << id
150
256
  end
151
257
  end
@@ -158,29 +264,40 @@ module RubyDecisionModel
158
264
  end
159
265
 
160
266
  if missing.any?
161
- raise MissingAnswers.new(
162
- "missing or wrong-type answers for: #{missing.join(', ')}",
163
- answers: normalized,
164
- missing: missing
165
- )
267
+ message = "missing or wrong-type answers for: #{missing.join(', ')}"
268
+ message += " (refused: #{refused.join(', ')})" if refused.any?
269
+ raise MissingAnswers.new(message, answers: normalized, missing: missing, refused: refused)
166
270
  end
167
271
 
168
272
  Response.new(
169
273
  answers: normalized,
170
- usage: normalize_usage(parsed["usage"]),
171
- model: parsed["model"],
172
- id: parsed["id"],
173
- raw: parsed
274
+ usage: @provider.usage(canonical),
275
+ model: canonical["model"],
276
+ id: canonical["id"],
277
+ raw: parsed,
278
+ request_id: request_id_from(response_headers)
174
279
  )
175
280
  end
176
281
 
282
+ def request_id_from(headers)
283
+ header = @provider.request_id_header
284
+ return nil unless header && headers.is_a?(Hash)
285
+
286
+ headers.each do |key, value|
287
+ next unless key.to_s.casecmp?(header)
288
+
289
+ return value.is_a?(Array) ? value.first : value
290
+ end
291
+ nil
292
+ end
293
+
177
294
  class MalformedAnswer < StandardError; end
178
295
 
179
296
  def normalize_answer(type, hash)
180
297
  case type
181
298
  when "noul"
182
299
  noul = hash["noul"]
183
- raise MalformedAnswer unless noul.is_a?(Numeric)
300
+ raise MalformedAnswer unless noul.is_a?(Numeric) && noul.between?(0, 1)
184
301
 
185
302
  Answers::Noul.new(noul: noul.to_f, probabilities: hash_or_empty(hash["probabilities"]))
186
303
  when "choice"
@@ -218,15 +335,5 @@ module RubyDecisionModel
218
335
  def hash_or_empty(value)
219
336
  value.is_a?(Hash) ? value : {}
220
337
  end
221
-
222
- def normalize_usage(usage)
223
- usage = {} unless usage.is_a?(Hash)
224
-
225
- Response::Usage.new(
226
- input_tokens: Integer(usage["input_tokens"], exception: false),
227
- output_tokens: Integer(usage["output_tokens"], exception: false),
228
- cost: Float(usage["cost"], exception: false)
229
- )
230
- end
231
338
  end
232
339
  end
@@ -19,12 +19,13 @@ module RubyDecisionModel
19
19
  class TimeoutError < TransportError; end
20
20
 
21
21
  class ApiError < Error
22
- attr_reader :status, :body
22
+ attr_reader :status, :body, :headers
23
23
 
24
- def initialize(message, status:, body:)
24
+ def initialize(message, status:, body:, headers: {})
25
25
  super(message)
26
26
  @status = status
27
27
  @body = body
28
+ @headers = headers || {}
28
29
  end
29
30
  end
30
31
 
@@ -32,8 +33,12 @@ module RubyDecisionModel
32
33
 
33
34
  class PayloadTooLarge < ApiError; end
34
35
 
36
+ class UnprocessableEntity < ApiError; end
37
+
35
38
  class RateLimited < ApiError; end
36
39
 
40
+ class Overloaded < ApiError; end
41
+
37
42
  class InvalidResponse < Error
38
43
  attr_reader :answers
39
44
 
@@ -43,12 +48,16 @@ module RubyDecisionModel
43
48
  end
44
49
  end
45
50
 
51
+ # Raised when any question went unanswered. `missing` lists every such id.
52
+ # `refused` lists the ones the provider explicitly declined, a subset of
53
+ # `missing`. Answers that did arrive are on `answers`.
46
54
  class MissingAnswers < InvalidResponse
47
- attr_reader :missing
55
+ attr_reader :missing, :refused
48
56
 
49
- def initialize(message, answers: {}, missing: [])
57
+ def initialize(message, answers: {}, missing: [], refused: [])
50
58
  super(message, answers: answers)
51
59
  @missing = missing
60
+ @refused = refused
52
61
  end
53
62
  end
54
63
  end
@@ -0,0 +1,33 @@
1
+ # frozen_string_literal: true
2
+
3
+ module RubyDecisionModel
4
+ # Builds the base64 data URLs that Client#ask takes as images:. Every
5
+ # provider that reads images wants them embedded; none fetch a remote URL.
6
+ module Images
7
+ CONTENT_TYPES = {
8
+ ".png" => "image/png",
9
+ ".jpg" => "image/jpeg",
10
+ ".jpeg" => "image/jpeg",
11
+ ".webp" => "image/webp",
12
+ ".gif" => "image/gif"
13
+ }.freeze
14
+
15
+ module_function
16
+
17
+ def data_url(bytes, content_type:)
18
+ raise ArgumentError, "content_type must be an image/* type, got #{content_type.inspect}" unless content_type.to_s.start_with?("image/")
19
+
20
+ # pack("m0") is strict base64 without the base64 gem, which leaves the
21
+ # default gems in Ruby 3.4.
22
+ "data:#{content_type};base64,#{[bytes].pack('m0')}"
23
+ end
24
+
25
+ # Reads a file and infers its type from the extension.
26
+ def from_file(path, content_type: nil)
27
+ content_type ||= CONTENT_TYPES[File.extname(path.to_s).downcase]
28
+ raise ArgumentError, "cannot infer an image type for #{path}; pass content_type:" if content_type.nil?
29
+
30
+ data_url(File.binread(path), content_type: content_type)
31
+ end
32
+ end
33
+ end