ruby_llm-contract 1.1.1 → 1.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.
Files changed (37) hide show
  1. checksums.yaml +4 -4
  2. data/CHANGELOG.md +103 -0
  3. data/README.md +18 -3
  4. data/docs/guide/getting_started.md +86 -1
  5. data/docs/guide/llm_judge.md +14 -0
  6. data/docs/guide/rails_integration.md +8 -0
  7. data/docs/guide/relation_to_agent.md +8 -2
  8. data/lib/ruby_llm/contract/adapters/response.rb +28 -2
  9. data/lib/ruby_llm/contract/adapters/ruby_llm.rb +136 -28
  10. data/lib/ruby_llm/contract/adapters/test.rb +6 -0
  11. data/lib/ruby_llm/contract/concerns/usage_aggregator.rb +8 -5
  12. data/lib/ruby_llm/contract/configuration.rb +5 -1
  13. data/lib/ruby_llm/contract/cost_calculator.rb +98 -44
  14. data/lib/ruby_llm/contract/eval/eval_history.rb +2 -1
  15. data/lib/ruby_llm/contract/eval/model_comparison.rb +11 -1
  16. data/lib/ruby_llm/contract/eval/recommender.rb +10 -3
  17. data/lib/ruby_llm/contract/eval/report_stats.rb +8 -3
  18. data/lib/ruby_llm/contract/eval/report_storage.rb +8 -2
  19. data/lib/ruby_llm/contract/pipeline/base.rb +3 -1
  20. data/lib/ruby_llm/contract/pipeline/runner.rb +11 -1
  21. data/lib/ruby_llm/contract/pipeline/trace.rb +10 -0
  22. data/lib/ruby_llm/contract/provider_options.rb +45 -0
  23. data/lib/ruby_llm/contract/step/adapter_caller.rb +2 -1
  24. data/lib/ruby_llm/contract/step/base.rb +29 -12
  25. data/lib/ruby_llm/contract/step/dsl.rb +54 -0
  26. data/lib/ruby_llm/contract/step/limit_checker.rb +48 -13
  27. data/lib/ruby_llm/contract/step/result_builder.rb +66 -4
  28. data/lib/ruby_llm/contract/step/retry_executor.rb +35 -4
  29. data/lib/ruby_llm/contract/step/retry_policy.rb +31 -5
  30. data/lib/ruby_llm/contract/step/runner.rb +20 -1
  31. data/lib/ruby_llm/contract/step/runner_config.rb +6 -3
  32. data/lib/ruby_llm/contract/step/trace.rb +61 -20
  33. data/lib/ruby_llm/contract/token_estimator.rb +5 -3
  34. data/lib/ruby_llm/contract/version.rb +1 -1
  35. data/lib/ruby_llm/contract/workflow_scope.rb +48 -0
  36. data/lib/ruby_llm/contract.rb +2 -0
  37. metadata +3 -1
@@ -26,9 +26,9 @@ module RubyLLM
26
26
  context: context).results.first
27
27
  end
28
28
 
29
- def estimate_cost(input:, model: nil, attachment: nil)
29
+ def estimate_cost(input:, model: nil, attachment: nil, provider: nil)
30
30
  model_name = estimated_model_name(model)
31
- model_info = CostCalculator.find_model(model_name)
31
+ model_info = CostCalculator.find_model(model_name, provider: provider)
32
32
  return nil unless model_info
33
33
 
34
34
  text_tokens = TokenEstimator.estimate(build_messages(input))
@@ -48,12 +48,13 @@ module RubyLLM
48
48
  output_tokens_estimate: output_tokens,
49
49
  estimated_cost: CostCalculator.calculate(
50
50
  model_name: model_name,
51
- usage: { input_tokens: input_tokens, output_tokens: output_tokens }
51
+ usage: { input_tokens: input_tokens, output_tokens: output_tokens },
52
+ provider: provider
52
53
  )
53
54
  }
54
55
  end
55
56
 
56
- def estimate_eval_cost(eval_name, models: nil)
57
+ def estimate_eval_cost(eval_name, models: nil, provider: nil)
57
58
  defn = send(:all_eval_definitions)[eval_name.to_s]
58
59
  raise ArgumentError, "No eval '#{eval_name}' defined" unless defn
59
60
 
@@ -61,7 +62,7 @@ module RubyLLM
61
62
  cases = defn.build_dataset.cases
62
63
 
63
64
  model_list.each_with_object({}) do |model_name, result|
64
- result[model_name] = estimate_eval_cost_for_model(cases, model_name)
65
+ result[model_name] = estimate_eval_cost_for_model(cases, model_name, provider)
65
66
  end
66
67
  end
67
68
 
@@ -92,7 +93,8 @@ module RubyLLM
92
93
 
93
94
  # Forwarded to the adapter as-is. A key known but not forwarded would be
94
95
  # silently dropped, so the known list is built from this one.
95
- ADAPTER_CONTEXT_KEYS = %i[provider assume_model_exists max_tokens reasoning_effort attachment].freeze
96
+ ADAPTER_CONTEXT_KEYS = %i[provider assume_model_exists max_tokens reasoning_effort attachment
97
+ provider_options].freeze
96
98
  KNOWN_CONTEXT_KEYS = (%i[adapter model temperature retry_policy_override] + ADAPTER_CONTEXT_KEYS).freeze
97
99
 
98
100
  include Concerns::ContextHelpers
@@ -101,7 +103,7 @@ module RubyLLM
101
103
  context = safe_context(context)
102
104
  warn_unknown_context_keys(context)
103
105
 
104
- result = dispatch_run(input, context)
106
+ result = WorkflowScope.workflow(name || "anonymous step") { dispatch_run(input, context) }
105
107
  log_result(result)
106
108
  invoke_around_call(input, result)
107
109
  end
@@ -145,9 +147,9 @@ module RubyLLM
145
147
  [0, true]
146
148
  end
147
149
 
148
- def estimate_eval_cost_for_model(cases, model_name)
150
+ def estimate_eval_cost_for_model(cases, model_name, provider)
149
151
  cases.sum do |test_case|
150
- estimate = estimate_cost(input: test_case.input, model: model_name)
152
+ estimate = estimate_cost(input: test_case.input, model: model_name, provider: provider)
151
153
  # Two misses floor to 0.0, not one: model absent from the registry
152
154
  # (nil estimate), or present with unreadable pricing (hash whose
153
155
  # estimated_cost is nil). Documented as a floor, not a fail-closed.
@@ -195,7 +197,7 @@ module RubyLLM
195
197
 
196
198
  def runtime_settings(context)
197
199
  policy = context.key?(:retry_policy_override) ? context[:retry_policy_override] : retry_policy
198
- extra = context.slice(*ADAPTER_CONTEXT_KEYS)
200
+ extra = context.slice(*ADAPTER_CONTEXT_KEYS).merge(merged_provider_options(context))
199
201
 
200
202
  # Always pass the class-level `thinking` config to the adapter when
201
203
  # set, so fields like `budget` survive a per-call `reasoning_effort`
@@ -220,6 +222,13 @@ module RubyLLM
220
222
  }
221
223
  end
222
224
 
225
+ # Class-level provider_options with the call's merged over them; a
226
+ # one-key hash so `merge` drops it entirely when neither is set.
227
+ def merged_provider_options(context)
228
+ merged = ProviderOptions.merge(provider_options, context[:provider_options])
229
+ merged ? { provider_options: merged } : {}
230
+ end
231
+
223
232
  def current_model_config
224
233
  policy = retry_policy
225
234
  if policy && policy.config_list.any?
@@ -274,7 +283,9 @@ module RubyLLM
274
283
  on_unknown_attachment_size: on_unknown_attachment_size,
275
284
  temperature: context_temperature || temperature,
276
285
  extra_options: extra_options,
277
- observers: class_observers
286
+ observers: class_observers,
287
+ on_incomplete_output: on_incomplete_output,
288
+ token_count: token_count
278
289
  )
279
290
  end
280
291
 
@@ -287,12 +298,18 @@ module RubyLLM
287
298
  "model=#{trace.model} status=#{result.status} " \
288
299
  "latency=#{trace.latency_ms}ms " \
289
300
  "tokens=#{trace.usage&.dig(:input_tokens) || 0}+#{trace.usage&.dig(:output_tokens) || 0} " \
290
- "cost=$#{format("%.6f", trace.cost || 0)}"
301
+ "cost=#{log_cost(trace)}"
291
302
  logger.info(msg)
292
303
 
293
304
  log_failed_observations(result, logger)
294
305
  end
295
306
 
307
+ def log_cost(trace)
308
+ return "unknown" if trace.respond_to?(:cost_unknown?) && trace.cost_unknown?
309
+
310
+ "$#{format("%.6f", trace.cost || 0)}"
311
+ end
312
+
296
313
  def log_failed_observations(result, logger)
297
314
  failed = result.observations.select { |o| !o[:passed] }
298
315
  return if failed.empty?
@@ -185,6 +185,46 @@ module RubyLLM
185
185
  inherited_value(:on_unknown_attachment_size) || UnknownPolicy::DEFAULT
186
186
  end
187
187
 
188
+ TOKEN_COUNT_MODES = %i[estimate exact].freeze
189
+ TOKEN_COUNT_DEFAULT = :estimate
190
+
191
+ # How max_input and max_cost measure the input before a call.
192
+ # `:estimate` (default) is the local chars/4 heuristic. `:exact` asks
193
+ # the provider (`chat.count_tokens`): one extra request per attempt,
194
+ # attachments included, no attachment_token_estimate needed. A provider
195
+ # or adapter that cannot count refuses the call (:limit_exceeded).
196
+ def token_count(mode = nil)
197
+ if mode
198
+ unless TOKEN_COUNT_MODES.include?(mode)
199
+ raise ArgumentError, "token_count must be one of #{TOKEN_COUNT_MODES.inspect}, got #{mode.inspect}"
200
+ end
201
+
202
+ return @token_count = mode
203
+ end
204
+
205
+ inherited_value(:token_count) || TOKEN_COUNT_DEFAULT
206
+ end
207
+
208
+ INCOMPLETE_OUTPUT_MODES = %i[accept refuse].freeze
209
+ INCOMPLETE_OUTPUT_DEFAULT = :accept
210
+
211
+ # `:refuse` fails a response the provider cut off at a token limit
212
+ # (:output_truncated) or stopped with a content filter
213
+ # (:content_filtered), before validation, so it is never :ok.
214
+ # `:accept` (default) validates it like any other response.
215
+ def on_incomplete_output(mode = nil)
216
+ if mode
217
+ unless INCOMPLETE_OUTPUT_MODES.include?(mode)
218
+ raise ArgumentError, "on_incomplete_output must be one of #{INCOMPLETE_OUTPUT_MODES.inspect}, " \
219
+ "got #{mode.inspect}"
220
+ end
221
+
222
+ return @on_incomplete_output = mode
223
+ end
224
+
225
+ inherited_value(:on_incomplete_output) || INCOMPLETE_OUTPUT_DEFAULT
226
+ end
227
+
188
228
  def model(name = nil)
189
229
  if name == :default
190
230
  @model = UNSET
@@ -215,6 +255,20 @@ module RubyLLM
215
255
  inherited_value_with_reset(:temperature)
216
256
  end
217
257
 
258
+ # Request options in the provider's vocabulary, e.g.
259
+ # `provider_options service_tier: "flex"`. Replaces an inherited hash;
260
+ # `provider_options :default` stops inheriting. `context: { provider_options: }`
261
+ # is merged over it per call.
262
+ def provider_options(options = nil)
263
+ if options == :default
264
+ @provider_options = UNSET
265
+ return nil
266
+ end
267
+ return @provider_options = ProviderOptions.frozen_copy(ProviderOptions.validate!(options)) if options
268
+
269
+ inherited_value_with_reset(:provider_options)
270
+ end
271
+
218
272
  def thinking(effort: nil, budget: nil)
219
273
  if effort == :default
220
274
  @thinking = UNSET
@@ -8,6 +8,7 @@ module RubyLLM
8
8
 
9
9
  def check_limits(messages)
10
10
  return nil unless max_input || max_cost
11
+ return check_counted_limits(messages) if exact_token_count?
11
12
 
12
13
  text_tokens = TokenEstimator.estimate(messages)
13
14
  attachment_tokens, attachment_error = resolve_attachment_tokens
@@ -23,6 +24,31 @@ module RubyLLM
23
24
  build_limit_result(messages, estimated, errors)
24
25
  end
25
26
 
27
+ # `token_count :exact`: the provider counts the input, attachments
28
+ # included. A count that cannot be had refuses the call rather than
29
+ # falling back to the heuristic under the name "exact".
30
+ def check_counted_limits(messages)
31
+ counted, count_error = count_input_tokens(messages)
32
+ return build_limit_result(messages, nil, [count_error], method: :exact) if count_error
33
+
34
+ errors = collect_limit_errors(counted, method: :exact)
35
+ errors.empty? ? nil : build_limit_result(messages, counted, errors, method: :exact)
36
+ end
37
+
38
+ def count_input_tokens(messages)
39
+ unless count_adapter.respond_to?(:count_tokens)
40
+ return [nil, "token_count :exact needs an adapter that counts tokens; " \
41
+ "#{count_adapter.class} does not"]
42
+ end
43
+
44
+ counted = count_adapter.count_tokens(messages: messages, **count_options)
45
+ return [counted, nil] if counted.is_a?(Integer) && !counted.negative?
46
+
47
+ [nil, "token_count :exact: no usable token count (got #{counted.inspect})"]
48
+ rescue ::RubyLLM::Error, ::Faraday::Error => e
49
+ [nil, "token_count :exact: the provider could not count the input tokens (#{e.message})"]
50
+ end
51
+
26
52
  # Fail-closed: when an attachment is passed via context but no
27
53
  # attachment_token_estimate is declared, the gem cannot bound vision/
28
54
  # PDF cost. Refuses with a clear error unless on_unknown_attachment_size
@@ -50,12 +76,17 @@ module RubyLLM
50
76
  [estimate, nil]
51
77
  end
52
78
 
53
- def collect_limit_errors(estimated)
79
+ def collect_limit_errors(estimated, method: :heuristic)
54
80
  errors = []
55
81
  if max_input && estimated > max_input
56
- errors << "Input token limit exceeded: estimated #{estimated} tokens (heuristic ±30%), max #{max_input}"
82
+ measured = if method == :exact
83
+ "#{estimated} tokens (counted by the provider)"
84
+ else
85
+ "estimated #{estimated} tokens (heuristic ±30%)"
86
+ end
87
+ errors << "Input token limit exceeded: #{measured}, max #{max_input}"
57
88
  end
58
- append_cost_error(estimated, errors) if max_cost
89
+ append_cost_error(estimated, errors, method) if max_cost
59
90
  errors
60
91
  end
61
92
 
@@ -66,18 +97,23 @@ module RubyLLM
66
97
  # for models expensive on completion side.
67
98
  DEFAULT_OUTPUT_RATIO = 1
68
99
 
69
- def append_cost_error(estimated, errors)
100
+ def append_cost_error(estimated, errors, method = :heuristic)
70
101
  estimated_output = effective_max_output || (estimated * DEFAULT_OUTPUT_RATIO)
71
102
  estimated_cost = CostCalculator.calculate(
72
103
  model_name: model_name,
73
- usage: { input_tokens: estimated, output_tokens: estimated_output }
104
+ usage: { input_tokens: estimated, output_tokens: estimated_output },
105
+ provider: pricing_provider
74
106
  )
75
107
 
76
108
  if estimated_cost.nil?
77
109
  handle_unknown_pricing(errors)
78
110
  elsif estimated_cost > max_cost
79
- errors << "Cost limit exceeded: estimated $#{format("%.6f", estimated_cost)} " \
80
- "(#{estimated} input + #{estimated_output} output tokens, heuristic ±30%), " \
111
+ tokens = if method == :exact
112
+ "#{estimated} input tokens counted by the provider + #{estimated_output} output"
113
+ else
114
+ "#{estimated} input + #{estimated_output} output tokens, heuristic ±30%"
115
+ end
116
+ errors << "Cost limit exceeded: estimated $#{format("%.6f", estimated_cost)} (#{tokens}), " \
81
117
  "max $#{format("%.6f", max_cost)}"
82
118
  end
83
119
  end
@@ -93,17 +129,16 @@ module RubyLLM
93
129
  end
94
130
  end
95
131
 
96
- def build_limit_result(messages, estimated, errors)
132
+ # No request was sent, so usage is zero; the measured input goes
133
+ # alongside, and is left out when it could not be measured.
134
+ def build_limit_result(messages, estimated, errors, method: :heuristic)
135
+ usage = { input_tokens: 0, output_tokens: 0, estimated_input_tokens: estimated, estimate_method: method }
97
136
  Result.new(
98
137
  status: :limit_exceeded,
99
138
  raw_output: nil,
100
139
  parsed_output: nil,
101
140
  validation_errors: errors,
102
- trace: Trace.new(
103
- messages: messages, model: model_name,
104
- usage: { input_tokens: 0, output_tokens: 0, estimated_input_tokens: estimated,
105
- estimate_method: :heuristic }
106
- )
141
+ trace: Trace.new(messages: messages, model: model_name, usage: usage.compact)
107
142
  )
108
143
  end
109
144
  end
@@ -4,12 +4,26 @@ module RubyLLM
4
4
  module Contract
5
5
  module Step
6
6
  class ResultBuilder
7
- def initialize(contract_definition:, output_type:, output_schema:, model:, observers:)
7
+ # Why the provider stopped, worded for a failed result. Anthropic also
8
+ # reports an exhausted context window as :max_tokens, so the message
9
+ # does not claim max_output was the limit hit.
10
+ FINISH_REASON_ERRORS = {
11
+ max_tokens: "provider reported token-limit termination",
12
+ content_filter: "provider reported content filtering or refusal"
13
+ }.freeze
14
+
15
+ # Statuses for `on_incomplete_output :refuse`.
16
+ INCOMPLETE_STATUSES = { max_tokens: :output_truncated, content_filter: :content_filtered }.freeze
17
+
18
+ def initialize(contract_definition:, output_type:, output_schema:, model:, observers:, max_output: nil,
19
+ on_incomplete_output: Dsl::INCOMPLETE_OUTPUT_DEFAULT)
8
20
  @contract_definition = contract_definition
9
21
  @output_type = output_type
10
22
  @output_schema = output_schema
11
23
  @model = model
12
24
  @observers = observers
25
+ @max_output = max_output
26
+ @on_incomplete_output = on_incomplete_output
13
27
  end
14
28
 
15
29
  def error_result(error_result:, messages:)
@@ -18,20 +32,24 @@ module RubyLLM
18
32
  raw_output: error_result.raw_output,
19
33
  parsed_output: error_result.parsed_output,
20
34
  validation_errors: error_result.validation_errors,
21
- trace: Trace.new(messages: messages, model: @model)
35
+ trace: Trace.new(messages: messages, model: @model, **error_fields(error_result.trace))
22
36
  )
23
37
  end
24
38
 
25
39
  def success_result(response:, messages:, latency_ms:, input:)
26
40
  raw_output = response.content
41
+ trace = Trace.new(messages: messages, model: @model, latency_ms: latency_ms, usage: response.usage,
42
+ **response_accounting(response))
43
+ incomplete = incomplete_result(response, trace)
44
+ return incomplete if incomplete
45
+
27
46
  validation_result = validate_output(raw_output, input)
28
- trace = Trace.new(messages: messages, model: @model, latency_ms: latency_ms, usage: response.usage)
29
47
 
30
48
  Result.new(
31
49
  status: validation_result[:status],
32
50
  raw_output: raw_output,
33
51
  parsed_output: validation_result[:parsed_output],
34
- validation_errors: validation_result[:errors],
52
+ validation_errors: validation_errors(validation_result, response),
35
53
  trace: trace,
36
54
  observations: observations_for(validation_result, input)
37
55
  )
@@ -39,6 +57,50 @@ module RubyLLM
39
57
 
40
58
  private
41
59
 
60
+ def error_fields(trace)
61
+ return {} unless trace.respond_to?(:error_class) && trace.error_class
62
+
63
+ { error_class: trace.error_class, error_type: trace.error_type }
64
+ end
65
+
66
+ # Decided before validation, so neither validate blocks nor observers
67
+ # see a response the step refuses.
68
+ def incomplete_result(response, trace)
69
+ return unless @on_incomplete_output == :refuse
70
+
71
+ status = INCOMPLETE_STATUSES[finish_reason(response)]
72
+ return unless status
73
+
74
+ Result.new(status: status, raw_output: response.content, parsed_output: nil,
75
+ validation_errors: [finish_reason_error(finish_reason(response))], trace: trace)
76
+ end
77
+
78
+ def finish_reason(response)
79
+ response.respond_to?(:finish_reason) ? response.finish_reason : nil
80
+ end
81
+
82
+ def finish_reason_error(reason)
83
+ message = FINISH_REASON_ERRORS[reason]
84
+ return message unless message && reason == :max_tokens && @max_output
85
+
86
+ "#{message}; configured max_output: #{@max_output}"
87
+ end
88
+
89
+ # Only a failed result says why the provider stopped; a result that
90
+ # passed validation keeps finish_reason in its trace alone.
91
+ def validation_errors(validation_result, response)
92
+ errors = validation_result[:errors]
93
+ return errors if validation_result[:status] == :ok
94
+
95
+ message = finish_reason_error(finish_reason(response))
96
+ message ? [*errors, message] : errors
97
+ end
98
+
99
+ # Custom adapters may return any object with `content` and `usage`.
100
+ def response_accounting(response)
101
+ response.respond_to?(:trace_accounting) ? response.trace_accounting : {}
102
+ end
103
+
42
104
  def observations_for(validation_result, input)
43
105
  return [] unless validation_result[:status] == :ok && @observers.any?
44
106
 
@@ -6,6 +6,9 @@ module RubyLLM
6
6
  module RetryExecutor
7
7
  include Concerns::UsageAggregator
8
8
 
9
+ # Copied from each attempt's trace into its attempts entry when present.
10
+ ATTEMPT_TRACE_FIELDS = %i[usage latency_ms cost finish_reason error_class].freeze
11
+
9
12
  private
10
13
 
11
14
  def run_with_retry(input, adapter:, default_model:, policy:, context_temperature: nil, extra_options: {})
@@ -39,7 +42,9 @@ module RubyLLM
39
42
  observations: last.observations,
40
43
  trace: last.trace.merge(
41
44
  attempts: attempt_log, usage: aggregated_usage,
42
- cost: total_cost, latency_ms: total_latency
45
+ cost: total_cost, latency_ms: total_latency,
46
+ usage_complete: aggregate_flag(all_attempts, :usage_complete),
47
+ cost_complete: aggregate_cost_complete(all_attempts)
43
48
  )
44
49
  )
45
50
  end
@@ -53,12 +58,38 @@ module RubyLLM
53
58
  end
54
59
 
55
60
  def append_trace_fields(entry, trace)
56
- entry[:usage] = trace.usage if trace.respond_to?(:usage) && trace.usage
57
- entry[:latency_ms] = trace.latency_ms if trace.respond_to?(:latency_ms) && trace.latency_ms
58
- entry[:cost] = trace.cost if trace.respond_to?(:cost) && trace.cost
61
+ ATTEMPT_TRACE_FIELDS.each do |field|
62
+ value = trace_value(trace, field)
63
+ entry[field] = value if value
64
+ end
65
+ entry[:cost_unknown] = true if trace_value(trace, :cost_unknown?)
66
+ entry[:usage_complete] = false if trace_value(trace, :usage_complete) == false
59
67
  entry
60
68
  end
61
69
 
70
+ def trace_value(trace, method)
71
+ trace.public_send(method) if trace.respond_to?(method)
72
+ end
73
+
74
+ # The subtotal is complete only if every attempt was; the last attempt's
75
+ # flag must not speak for it. nil when no attempt reported accounting
76
+ # state, so a retried Test-adapter step serializes exactly as before.
77
+ def aggregate_cost_complete(all_attempts)
78
+ traces = all_attempts.map { |a| a[:result].trace }
79
+ return false if traces.any? { |trace| trace.respond_to?(:cost_unknown?) && trace.cost_unknown? }
80
+
81
+ aggregate_flag(all_attempts, :cost_complete)
82
+ end
83
+
84
+ def aggregate_flag(all_attempts, flag)
85
+ flags = all_attempts.map { |a| a[:result].trace }
86
+ .select { |trace| trace.respond_to?(flag) }
87
+ .map { |trace| trace.public_send(flag) }.compact
88
+ return nil if flags.empty?
89
+
90
+ flags.all?
91
+ end
92
+
62
93
  def sum_attempt_costs(all_attempts)
63
94
  costs = extract_trace_values(all_attempts, :cost)
64
95
  return nil if costs.empty?
@@ -4,7 +4,7 @@ module RubyLLM
4
4
  module Contract
5
5
  module Step
6
6
  class RetryPolicy
7
- attr_reader :max_attempts, :retryable_statuses
7
+ attr_reader :max_attempts, :retryable_statuses, :retryable_errors
8
8
 
9
9
  # Breaking (0.7.0): :adapter_error removed from defaults. ruby_llm's Faraday
10
10
  # middleware already retries transport errors (RateLimitError, ServerError,
@@ -18,6 +18,7 @@ module RubyLLM
18
18
  def initialize(models: nil, attempts: nil, retry_on: nil, &block)
19
19
  @configs = []
20
20
  @retryable_statuses = DEFAULT_RETRY_ON.dup
21
+ @retryable_errors = []
21
22
 
22
23
  if block
23
24
  @max_attempts = 1
@@ -49,12 +50,20 @@ module RubyLLM
49
50
  @configs
50
51
  end
51
52
 
52
- def retry_on(*statuses)
53
- @retryable_statuses = statuses.flatten
53
+ # Statuses, and exception classes for adapter errors:
54
+ # `retry_on :validation_failed, RubyLLM::RateLimitError` retries an
55
+ # adapter error only when it was a RateLimitError (or a subclass), while
56
+ # `:adapter_error` retries any. Replaces the defaults, as before.
57
+ def retry_on(*conditions)
58
+ @retryable_statuses, @retryable_errors = split_conditions(conditions.flatten)
54
59
  end
55
60
 
56
61
  def retryable?(result)
57
- retryable_statuses.include?(result.status)
62
+ return true if retryable_statuses.include?(result.status)
63
+ return false unless result.status == :adapter_error && retryable_errors.any?
64
+
65
+ error_type = result.trace.respond_to?(:error_type) ? result.trace.error_type : nil
66
+ !error_type.nil? && retryable_errors.any? { |klass| error_type <= klass }
58
67
  end
59
68
 
60
69
  def config_for_attempt(attempt, default_config)
@@ -78,7 +87,24 @@ module RubyLLM
78
87
  else
79
88
  @max_attempts = attempts || 1
80
89
  end
81
- @retryable_statuses = Array(retry_on).dup if retry_on
90
+ @retryable_statuses, @retryable_errors = split_conditions(Array(retry_on)) if retry_on
91
+ end
92
+
93
+ # Symbols are statuses; exception classes match adapter errors. Anything
94
+ # else (a String status among them) would never match, so it raises
95
+ # instead of quietly turning retries off.
96
+ def split_conditions(conditions)
97
+ invalid = conditions.reject { |condition| condition.is_a?(Symbol) || exception_class?(condition) }
98
+ unless invalid.empty?
99
+ raise ArgumentError,
100
+ "retry_on takes status symbols and exception classes, got #{invalid.map(&:inspect).join(", ")}"
101
+ end
102
+
103
+ conditions.partition { |condition| condition.is_a?(Symbol) }
104
+ end
105
+
106
+ def exception_class?(condition)
107
+ condition.is_a?(Class) && condition <= Exception
82
108
  end
83
109
 
84
110
  def normalize_config(entry)
@@ -59,7 +59,9 @@ module RubyLLM
59
59
  output_type: @config.output_type,
60
60
  output_schema: @config.output_schema,
61
61
  model: @config.model,
62
- observers: @config.observers
62
+ observers: @config.observers,
63
+ max_output: @config.effective_max_output,
64
+ on_incomplete_output: @config.on_incomplete_output
63
65
  )
64
66
  end
65
67
 
@@ -83,10 +85,27 @@ module RubyLLM
83
85
  @config.attachment_token_estimate
84
86
  end
85
87
 
88
+ def exact_token_count?
89
+ @config.token_count == :exact
90
+ end
91
+
92
+ def count_adapter
93
+ @config.adapter
94
+ end
95
+
96
+ def count_options
97
+ @config.adapter_options
98
+ end
99
+
86
100
  def on_unknown_attachment_size
87
101
  @config.on_unknown_attachment_size
88
102
  end
89
103
 
104
+ # The provider the adapter will call, so the estimate uses its price.
105
+ def pricing_provider
106
+ @config.extra_options&.dig(:provider)
107
+ end
108
+
90
109
  def attachment_present?
91
110
  opts = @config.extra_options
92
111
  !opts.nil? && !opts[:attachment].nil?
@@ -19,7 +19,9 @@ module RubyLLM
19
19
  :on_unknown_attachment_size,
20
20
  :temperature,
21
21
  :extra_options,
22
- :observers
22
+ :observers,
23
+ :on_incomplete_output,
24
+ :token_count
23
25
  ) do
24
26
  # Factory with sensible defaults for optional fields. Lets callers
25
27
  # (Step::Base#run_once and tests) construct a RunnerConfig without
@@ -30,7 +32,8 @@ module RubyLLM
30
32
  output_schema: nil, max_output: nil,
31
33
  max_input: nil, max_cost: nil, on_unknown_pricing: UnknownPolicy::DEFAULT,
32
34
  attachment_token_estimate: nil, on_unknown_attachment_size: UnknownPolicy::DEFAULT,
33
- temperature: nil, extra_options: {}, observers: [])
35
+ temperature: nil, extra_options: {}, observers: [],
36
+ on_incomplete_output: Dsl::INCOMPLETE_OUTPUT_DEFAULT, token_count: Dsl::TOKEN_COUNT_DEFAULT)
34
37
  new(
35
38
  input_type: input_type, output_type: output_type,
36
39
  prompt_block: prompt_block, contract_definition: contract_definition,
@@ -41,7 +44,7 @@ module RubyLLM
41
44
  attachment_token_estimate: attachment_token_estimate,
42
45
  on_unknown_attachment_size: on_unknown_attachment_size,
43
46
  temperature: temperature, extra_options: extra_options,
44
- observers: observers
47
+ observers: observers, on_incomplete_output: on_incomplete_output, token_count: token_count
45
48
  )
46
49
  end
47
50