roast-ai 1.0.2 → 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.
- checksums.yaml +4 -4
- data/.claude/commands/docs/write-comments.md +1 -1
- data/.rubocop.yml +12 -1
- data/Gemfile +2 -2
- data/Gemfile.lock +149 -34
- data/README.md +56 -3
- data/examples/agent_with_multiple_prompts.rb +27 -0
- data/examples/custom_logging.rb +4 -2
- data/examples/demo/Gemfile.lock +49 -15
- data/examples/plugin-gem-example/Gemfile.lock +19 -15
- data/examples/simple_chat.rb +1 -1
- data/examples/simple_pi_agent.rb +18 -0
- data/internal/rubocop/cop/roast/no_test_class_nesting.rb +126 -0
- data/internal/rubocop/rubocop-roast.yml +6 -0
- data/internal/workflows/maintenance/branch_docs_impact.rb +97 -0
- data/internal/workflows/maintenance/deprecated_models_docs_updater.rb +78 -0
- data/lib/roast/cog/config.rb +1 -1
- data/lib/roast/cog/output.rb +2 -1
- data/lib/roast/cog/registry.rb +3 -3
- data/lib/roast/cog_input_manager.rb +28 -7
- data/lib/roast/cogs/agent/config.rb +2 -2
- data/lib/roast/cogs/agent/input.rb +20 -22
- data/lib/roast/cogs/agent/providers/claude/claude_invocation.rb +13 -5
- data/lib/roast/cogs/agent/providers/claude/messages/result_message.rb +1 -1
- data/lib/roast/cogs/agent/providers/claude/tool_result.rb +344 -4
- data/lib/roast/cogs/agent/providers/claude/tool_use.rb +356 -1
- data/lib/roast/cogs/agent/providers/claude.rb +16 -3
- data/lib/roast/cogs/agent/providers/pi/messages/tool_call_message.rb +60 -0
- data/lib/roast/cogs/agent/providers/pi/messages/tool_result_message.rb +57 -0
- data/lib/roast/cogs/agent/providers/pi/pi_invocation.rb +352 -0
- data/lib/roast/cogs/agent/providers/pi.rb +41 -0
- data/lib/roast/cogs/agent/stats.rb +29 -0
- data/lib/roast/cogs/agent/usage.rb +22 -0
- data/lib/roast/cogs/agent.rb +5 -6
- data/lib/roast/cogs/chat/config.rb +28 -2
- data/lib/roast/cogs/chat.rb +82 -10
- data/lib/roast/event.rb +1 -0
- data/lib/roast/event_monitor.rb +35 -3
- data/lib/roast/log.rb +21 -0
- data/lib/roast/log_formatter.rb +9 -7
- data/lib/roast/version.rb +1 -1
- data/lib/roast.rb +1 -3
- data/roast-ai.gemspec +2 -1
- data/sorbet/rbi/gems/activesupport@8.0.2.rbi +549 -383
- data/sorbet/rbi/gems/addressable@2.8.7.rbi +46 -44
- data/sorbet/rbi/gems/ast@2.4.3.rbi +7 -6
- data/sorbet/rbi/gems/async@2.34.0.rbi +21 -3
- data/sorbet/rbi/gems/benchmark@0.4.1.rbi +7 -7
- data/sorbet/rbi/gems/bigdecimal@3.2.2.rbi +198 -1
- data/sorbet/rbi/gems/concurrent-ruby@1.3.5.rbi +405 -328
- data/sorbet/rbi/gems/console@1.34.2.rbi +2 -2
- data/sorbet/rbi/gems/docile@1.4.1.rbi +30 -30
- data/sorbet/rbi/gems/drb@2.2.3.rbi +25 -25
- data/sorbet/rbi/gems/erubi@1.13.1.rbi +2 -0
- data/sorbet/rbi/gems/faraday-net_http@3.4.2.rbi +2 -77
- data/sorbet/rbi/gems/faraday-retry@2.3.2.rbi +2 -57
- data/sorbet/rbi/gems/faraday@2.14.1.rbi +382 -75
- data/sorbet/rbi/gems/guard-compat@1.2.1.rbi +1 -110
- data/sorbet/rbi/gems/guard-minitest@2.4.6.rbi +0 -139
- data/sorbet/rbi/gems/guard@2.19.1.rbi +38 -38
- data/sorbet/rbi/gems/hashdiff@1.2.0.rbi +3 -3
- data/sorbet/rbi/gems/i18n@1.14.7.rbi +53 -29
- data/sorbet/rbi/gems/io-event@1.14.0.rbi +67 -10
- data/sorbet/rbi/gems/json@2.18.1.rbi +227 -5
- data/sorbet/rbi/gems/lint_roller@1.1.0.rbi +83 -0
- data/sorbet/rbi/gems/listen@3.9.0.rbi +7 -7
- data/sorbet/rbi/gems/logger@1.7.0.rbi +3 -3
- data/sorbet/rbi/gems/lumberjack@1.2.10.rbi +21 -21
- data/sorbet/rbi/gems/marcel@1.1.0.rbi +1 -1
- data/sorbet/rbi/gems/minitest-rg@5.3.0.rbi +0 -96
- data/sorbet/rbi/gems/minitest@5.25.5.rbi +1 -16
- data/sorbet/rbi/gems/net-http@0.9.1.rbi +27 -19
- data/sorbet/rbi/gems/netrc@0.11.0.rbi +18 -0
- data/sorbet/rbi/gems/notiffany@0.1.3.rbi +20 -20
- data/sorbet/rbi/gems/ostruct@0.6.2.rbi +149 -15
- data/sorbet/rbi/gems/parser@3.3.8.0.rbi +141 -139
- data/sorbet/rbi/gems/prism@1.4.0.rbi +922 -864
- data/sorbet/rbi/gems/public_suffix@6.0.2.rbi +56 -35
- data/sorbet/rbi/gems/racc@1.8.1.rbi +10 -2
- data/sorbet/rbi/gems/rainbow@3.1.1.rbi +12 -12
- data/sorbet/rbi/gems/rake@13.3.0.rbi +219 -318
- data/sorbet/rbi/gems/{rbi@0.3.6.rbi → rbi@0.3.9.rbi} +612 -2267
- data/sorbet/rbi/gems/{rbs@3.9.4.rbi → rbs@4.0.0.dev.5.rbi} +2013 -680
- data/sorbet/rbi/gems/regexp_parser@2.10.0.rbi +151 -113
- data/sorbet/rbi/gems/require-hooks@0.2.3.rbi +110 -0
- data/sorbet/rbi/gems/rexml@3.4.2.rbi +24 -51
- data/sorbet/rbi/gems/rubocop-ast@1.45.1.rbi +506 -815
- data/sorbet/rbi/gems/rubocop-sorbet@0.10.5.rbi +16 -16
- data/sorbet/rbi/gems/rubocop@1.77.0.rbi +2692 -2327
- data/sorbet/rbi/gems/ruby-progressbar@1.13.0.rbi +8 -8
- data/sorbet/rbi/gems/ruby_llm@1.8.2.rbi +38 -23
- data/sorbet/rbi/gems/securerandom@0.4.1.rbi +1 -1
- data/sorbet/rbi/gems/simplecov-html@0.13.2.rbi +2 -131
- data/sorbet/rbi/gems/simplecov@0.22.0.rbi +28 -127
- data/sorbet/rbi/gems/{spoom@1.6.3.rbi → spoom@1.7.11.rbi} +1139 -2246
- data/sorbet/rbi/gems/sqlite3@2.9.0.rbi +91 -1
- data/sorbet/rbi/gems/{tapioca@0.16.11.rbi → tapioca@0.17.10.rbi} +721 -835
- data/sorbet/rbi/gems/thor@1.4.0.rbi +53 -53
- data/sorbet/rbi/gems/tsort@0.2.0.rbi +393 -0
- data/sorbet/rbi/gems/type_toolkit@0.0.5.rbi +49 -0
- data/sorbet/rbi/gems/tzinfo@2.0.6.rbi +144 -143
- data/sorbet/rbi/gems/uri@1.1.1.rbi +7 -7
- data/sorbet/rbi/gems/vcr@6.3.1.rbi +53 -36
- data/sorbet/rbi/gems/webmock@3.25.1.rbi +38 -13
- data/sorbet/rbi/gems/zeitwerk@2.7.3.rbi +39 -272
- data/sorbet/rbi/shims/lib/roast/execution_context.rbi +3 -3
- data/tutorial/01_your_first_workflow/README.md +9 -5
- data/tutorial/01_your_first_workflow/configured_chat.rb +1 -1
- data/tutorial/02_chaining_cogs/README.md +2 -2
- data/tutorial/02_chaining_cogs/code_review.rb +1 -1
- data/tutorial/02_chaining_cogs/session_resumption.rb +1 -1
- data/tutorial/03_targets_and_params/README.md +1 -1
- data/tutorial/04_configuration_options/README.md +2 -2
- data/tutorial/08_iterative_workflows/README.md +1 -1
- data/tutorial/README.md +1 -1
- metadata +39 -17
- data/docs/AGENT_STEPS.md +0 -288
- data/docs/INSTRUMENTATION.md +0 -243
- data/docs/ITERATION_SYNTAX.md +0 -147
- data/docs/VALIDATION.md +0 -178
- data/lib/roast/nil_assertions.rb +0 -23
- /data/internal/documentation/{architectural-notes.md → comments/architectural-notes.md} +0 -0
- /data/internal/documentation/{doc-comments-external.md → comments/doc-comments-external.md} +0 -0
- /data/internal/documentation/{doc-comments-internal.md → comments/doc-comments-internal.md} +0 -0
- /data/internal/documentation/{doc-comments.md → comments/doc-comments.md} +0 -0
|
@@ -0,0 +1,352 @@
|
|
|
1
|
+
# typed: true
|
|
2
|
+
# frozen_string_literal: true
|
|
3
|
+
|
|
4
|
+
module Roast
|
|
5
|
+
module Cogs
|
|
6
|
+
class Agent < Cog
|
|
7
|
+
module Providers
|
|
8
|
+
class Pi < Provider
|
|
9
|
+
class PiInvocation
|
|
10
|
+
class PiInvocationError < Roast::Error; end
|
|
11
|
+
|
|
12
|
+
class PiNotStartedError < PiInvocationError; end
|
|
13
|
+
|
|
14
|
+
class PiAlreadyStartedError < PiInvocationError; end
|
|
15
|
+
|
|
16
|
+
class PiNotCompletedError < PiInvocationError; end
|
|
17
|
+
|
|
18
|
+
class PiFailedError < PiInvocationError; end
|
|
19
|
+
|
|
20
|
+
class Context
|
|
21
|
+
def initialize
|
|
22
|
+
@tool_calls = {} #: Hash[String, Messages::ToolCallMessage]
|
|
23
|
+
end
|
|
24
|
+
|
|
25
|
+
#: (String?) -> Messages::ToolCallMessage?
|
|
26
|
+
def tool_call(tool_call_id)
|
|
27
|
+
@tool_calls[tool_call_id] if tool_call_id
|
|
28
|
+
end
|
|
29
|
+
|
|
30
|
+
#: (Messages::ToolCallMessage) -> void
|
|
31
|
+
def add_tool_call(tool_call_message)
|
|
32
|
+
id = tool_call_message.id
|
|
33
|
+
@tool_calls[id] = tool_call_message if id
|
|
34
|
+
end
|
|
35
|
+
end
|
|
36
|
+
|
|
37
|
+
class Result
|
|
38
|
+
#: String
|
|
39
|
+
attr_accessor :response
|
|
40
|
+
|
|
41
|
+
#: bool
|
|
42
|
+
attr_accessor :success
|
|
43
|
+
|
|
44
|
+
#: String?
|
|
45
|
+
attr_accessor :session
|
|
46
|
+
|
|
47
|
+
#: Stats?
|
|
48
|
+
attr_accessor :stats
|
|
49
|
+
|
|
50
|
+
def initialize
|
|
51
|
+
@response = ""
|
|
52
|
+
@success = false
|
|
53
|
+
end
|
|
54
|
+
end
|
|
55
|
+
|
|
56
|
+
#: (Agent::Config, String, String?) -> void
|
|
57
|
+
def initialize(config, prompt, session)
|
|
58
|
+
@base_command = config.valid_command #: (String | Array[String])?
|
|
59
|
+
@model = config.valid_model #: String?
|
|
60
|
+
@append_system_prompt = config.valid_append_system_prompt #: String?
|
|
61
|
+
@replace_system_prompt = config.valid_replace_system_prompt #: String?
|
|
62
|
+
@working_directory = config.valid_working_directory #: Pathname?
|
|
63
|
+
@prompt = prompt #: String
|
|
64
|
+
@session = session #: String?
|
|
65
|
+
@context = Context.new #: Context
|
|
66
|
+
@result = Result.new #: Result
|
|
67
|
+
@raw_dump_file = config.valid_dump_raw_agent_messages_to_path #: Pathname?
|
|
68
|
+
@show_prompt = config.show_prompt? #: bool
|
|
69
|
+
@show_progress = config.show_progress? #: bool
|
|
70
|
+
@show_response = config.show_response? #: bool
|
|
71
|
+
@num_turns = 0 #: Integer
|
|
72
|
+
@total_cost = 0.0 #: Float
|
|
73
|
+
@model_usage_accumulator = {} #: Hash[String, Hash[Symbol, Numeric]]
|
|
74
|
+
@current_text_block = +"" #: String
|
|
75
|
+
@start_time_ms = nil #: Integer?
|
|
76
|
+
end
|
|
77
|
+
|
|
78
|
+
#: () -> void
|
|
79
|
+
def run!
|
|
80
|
+
raise PiAlreadyStartedError if started?
|
|
81
|
+
|
|
82
|
+
@started = true
|
|
83
|
+
Event << { block: { header: "USER PROMPT", content: @prompt } } if @show_prompt
|
|
84
|
+
@start_time_ms = (Process.clock_gettime(Process::CLOCK_MONOTONIC) * 1000).to_i
|
|
85
|
+
_stdout, stderr, status = CommandRunner.execute(
|
|
86
|
+
command_line,
|
|
87
|
+
working_directory: @working_directory,
|
|
88
|
+
stdin_content: @prompt,
|
|
89
|
+
stdout_handler: lambda { |line| handle_stdout(line) },
|
|
90
|
+
)
|
|
91
|
+
@end_time_ms = (Process.clock_gettime(Process::CLOCK_MONOTONIC) * 1000).to_i #: Integer?
|
|
92
|
+
|
|
93
|
+
if status.success?
|
|
94
|
+
@completed = true
|
|
95
|
+
@result.success = true
|
|
96
|
+
finalize_stats!
|
|
97
|
+
Event << { block: { header: "AGENT RESPONSE", content: @result.response } } if @show_response
|
|
98
|
+
else
|
|
99
|
+
@failed = true
|
|
100
|
+
@result.success = false
|
|
101
|
+
@result.response += "\n" unless @result.response.blank? || @result.response.ends_with?("\n")
|
|
102
|
+
@result.response += stderr
|
|
103
|
+
end
|
|
104
|
+
end
|
|
105
|
+
|
|
106
|
+
#: () -> bool
|
|
107
|
+
def started?
|
|
108
|
+
@started ||= false
|
|
109
|
+
end
|
|
110
|
+
|
|
111
|
+
#: () -> bool
|
|
112
|
+
def running?
|
|
113
|
+
started? && !completed? && !failed?
|
|
114
|
+
end
|
|
115
|
+
|
|
116
|
+
#: () -> bool
|
|
117
|
+
def completed?
|
|
118
|
+
@completed ||= false
|
|
119
|
+
end
|
|
120
|
+
|
|
121
|
+
#: () -> bool
|
|
122
|
+
def failed?
|
|
123
|
+
@failed ||= false
|
|
124
|
+
end
|
|
125
|
+
|
|
126
|
+
#: () -> Result
|
|
127
|
+
def result
|
|
128
|
+
raise PiNotStartedError unless started?
|
|
129
|
+
raise PiFailedError, @result.response if failed?
|
|
130
|
+
raise PiNotCompletedError, @result.response unless completed?
|
|
131
|
+
|
|
132
|
+
@result
|
|
133
|
+
end
|
|
134
|
+
|
|
135
|
+
private
|
|
136
|
+
|
|
137
|
+
#: (String) -> void
|
|
138
|
+
def handle_stdout(line)
|
|
139
|
+
line = line.strip
|
|
140
|
+
return if line.empty?
|
|
141
|
+
|
|
142
|
+
if @raw_dump_file
|
|
143
|
+
@raw_dump_file.dirname.mkpath
|
|
144
|
+
File.write(@raw_dump_file.to_s, "#{line}\n", mode: "a")
|
|
145
|
+
end
|
|
146
|
+
|
|
147
|
+
begin
|
|
148
|
+
data = JSON.parse(line, symbolize_names: true)
|
|
149
|
+
rescue JSON::ParserError
|
|
150
|
+
return
|
|
151
|
+
end
|
|
152
|
+
|
|
153
|
+
handle_message(data)
|
|
154
|
+
end
|
|
155
|
+
|
|
156
|
+
#: (Hash[Symbol, untyped]) -> void
|
|
157
|
+
def handle_message(data)
|
|
158
|
+
type = data[:type]&.to_sym
|
|
159
|
+
|
|
160
|
+
case type
|
|
161
|
+
when :session
|
|
162
|
+
handle_session(data)
|
|
163
|
+
when :turn_start
|
|
164
|
+
@num_turns += 1
|
|
165
|
+
when :turn_end
|
|
166
|
+
# turn_end contains the final assistant message and tool results for this turn
|
|
167
|
+
when :message_update
|
|
168
|
+
handle_message_update(data)
|
|
169
|
+
when :message_end
|
|
170
|
+
handle_message_end(data)
|
|
171
|
+
when :tool_execution_start
|
|
172
|
+
handle_tool_execution_start(data)
|
|
173
|
+
when :tool_execution_end
|
|
174
|
+
handle_tool_execution_end(data)
|
|
175
|
+
when :agent_end
|
|
176
|
+
handle_agent_end(data)
|
|
177
|
+
when :agent_start, :message_start, :tool_execution_update
|
|
178
|
+
# These are informational; no action needed
|
|
179
|
+
end
|
|
180
|
+
end
|
|
181
|
+
|
|
182
|
+
#: (Hash[Symbol, untyped]) -> void
|
|
183
|
+
def handle_session(data)
|
|
184
|
+
session_id = data[:id]
|
|
185
|
+
if session_id.present? && @result.session != session_id
|
|
186
|
+
Event << { debug: "New Pi Session ID: #{session_id}" }
|
|
187
|
+
@result.session = session_id
|
|
188
|
+
end
|
|
189
|
+
end
|
|
190
|
+
|
|
191
|
+
#: (Hash[Symbol, untyped]) -> void
|
|
192
|
+
def handle_message_update(data)
|
|
193
|
+
event = data[:assistantMessageEvent]
|
|
194
|
+
return unless event
|
|
195
|
+
|
|
196
|
+
event_type = event[:type]&.to_sym
|
|
197
|
+
|
|
198
|
+
case event_type
|
|
199
|
+
when :text_delta
|
|
200
|
+
delta = event[:delta]
|
|
201
|
+
@current_text_block << delta if delta
|
|
202
|
+
when :text_end
|
|
203
|
+
content = event[:content]
|
|
204
|
+
if content.present?
|
|
205
|
+
@result.response = content
|
|
206
|
+
elsif @current_text_block.present?
|
|
207
|
+
@result.response = @current_text_block.dup
|
|
208
|
+
end
|
|
209
|
+
# Print the accumulated text block as a single unit (like Claude does)
|
|
210
|
+
puts @current_text_block if @current_text_block.present? && @show_progress
|
|
211
|
+
@current_text_block = +""
|
|
212
|
+
when :toolcall_end
|
|
213
|
+
tool_call = event[:toolCall]
|
|
214
|
+
if tool_call
|
|
215
|
+
tool_call_msg = Messages::ToolCallMessage.new(
|
|
216
|
+
id: tool_call[:id],
|
|
217
|
+
name: tool_call[:name],
|
|
218
|
+
arguments: tool_call[:arguments] || {},
|
|
219
|
+
)
|
|
220
|
+
@context.add_tool_call(tool_call_msg)
|
|
221
|
+
end
|
|
222
|
+
end
|
|
223
|
+
end
|
|
224
|
+
|
|
225
|
+
#: (Hash[Symbol, untyped]) -> void
|
|
226
|
+
def handle_message_end(data)
|
|
227
|
+
message = data[:message]
|
|
228
|
+
return unless message
|
|
229
|
+
|
|
230
|
+
role = message[:role]&.to_sym
|
|
231
|
+
|
|
232
|
+
case role
|
|
233
|
+
when :assistant
|
|
234
|
+
# Extract usage from the final assistant message_end
|
|
235
|
+
usage = message[:usage]
|
|
236
|
+
model = message[:model]
|
|
237
|
+
if usage && model
|
|
238
|
+
accumulate_usage(model, usage)
|
|
239
|
+
end
|
|
240
|
+
|
|
241
|
+
# Extract final text if present
|
|
242
|
+
content = message[:content]
|
|
243
|
+
if content.is_a?(Array)
|
|
244
|
+
text_parts = content.select { |c| c[:type] == "text" }.map { |c| c[:text] }
|
|
245
|
+
@result.response = text_parts.join if text_parts.any?
|
|
246
|
+
end
|
|
247
|
+
end
|
|
248
|
+
end
|
|
249
|
+
|
|
250
|
+
#: (Hash[Symbol, untyped]) -> void
|
|
251
|
+
def handle_tool_execution_start(data)
|
|
252
|
+
tool_name = data[:toolName]
|
|
253
|
+
args = data[:args]
|
|
254
|
+
return unless @show_progress && tool_name
|
|
255
|
+
|
|
256
|
+
formatted = Messages::ToolCallMessage.new(
|
|
257
|
+
id: data[:toolCallId],
|
|
258
|
+
name: tool_name,
|
|
259
|
+
arguments: args || {},
|
|
260
|
+
).format
|
|
261
|
+
puts formatted if formatted.present?
|
|
262
|
+
end
|
|
263
|
+
|
|
264
|
+
#: (Hash[Symbol, untyped]) -> void
|
|
265
|
+
def handle_tool_execution_end(data)
|
|
266
|
+
return unless @show_progress
|
|
267
|
+
|
|
268
|
+
result_data = data[:result]
|
|
269
|
+
content = result_data&.dig(:content)&.first&.dig(:text) if result_data
|
|
270
|
+
formatted = Messages::ToolResultMessage.new(
|
|
271
|
+
tool_call_id: data[:toolCallId],
|
|
272
|
+
tool_name: data[:toolName],
|
|
273
|
+
content: content,
|
|
274
|
+
is_error: data[:isError] || false,
|
|
275
|
+
).format(@context)
|
|
276
|
+
puts formatted if formatted.present?
|
|
277
|
+
end
|
|
278
|
+
|
|
279
|
+
#: (Hash[Symbol, untyped]) -> void
|
|
280
|
+
def handle_agent_end(data)
|
|
281
|
+
# Extract final response from the last assistant message
|
|
282
|
+
messages = data[:messages]
|
|
283
|
+
return unless messages.is_a?(Array)
|
|
284
|
+
|
|
285
|
+
last_assistant = messages.reverse.find { |m| m[:role] == "assistant" }
|
|
286
|
+
return unless last_assistant
|
|
287
|
+
|
|
288
|
+
content = last_assistant[:content]
|
|
289
|
+
if content.is_a?(Array)
|
|
290
|
+
text_parts = content.select { |c| c[:type] == "text" }.map { |c| c[:text] }
|
|
291
|
+
@result.response = text_parts.join if text_parts.any?
|
|
292
|
+
end
|
|
293
|
+
end
|
|
294
|
+
|
|
295
|
+
#: (String, Hash[Symbol, untyped]) -> void
|
|
296
|
+
def accumulate_usage(model, usage)
|
|
297
|
+
acc = @model_usage_accumulator[model] ||= { input: 0, output: 0, cache_read: 0, cache_write: 0, cost: 0.0 }
|
|
298
|
+
acc[:input] = (acc[:input] || 0) + (usage[:input] || 0)
|
|
299
|
+
acc[:output] = (acc[:output] || 0) + (usage[:output] || 0)
|
|
300
|
+
acc[:cache_read] = (acc[:cache_read] || 0) + (usage[:cacheRead] || 0)
|
|
301
|
+
acc[:cache_write] = (acc[:cache_write] || 0) + (usage[:cacheWrite] || 0)
|
|
302
|
+
cost = usage.dig(:cost, :total) || 0.0
|
|
303
|
+
acc[:cost] = (acc[:cost] || 0.0) + cost
|
|
304
|
+
@total_cost = @model_usage_accumulator.values.sum(0.0) { |a| a[:cost].to_f }
|
|
305
|
+
end
|
|
306
|
+
|
|
307
|
+
#: () -> void
|
|
308
|
+
def finalize_stats!
|
|
309
|
+
stats = Stats.new
|
|
310
|
+
stats.num_turns = @num_turns
|
|
311
|
+
stats.duration_ms = @end_time_ms - @start_time_ms if @start_time_ms && @end_time_ms
|
|
312
|
+
|
|
313
|
+
@model_usage_accumulator.each do |model, acc|
|
|
314
|
+
usage = Usage.new
|
|
315
|
+
usage.input_tokens = acc[:input].to_i
|
|
316
|
+
usage.output_tokens = acc[:output].to_i
|
|
317
|
+
usage.cost_usd = acc[:cost].to_f
|
|
318
|
+
stats.model_usage[model] = usage
|
|
319
|
+
stats.usage.input_tokens = (stats.usage.input_tokens || 0) + (usage.input_tokens || 0)
|
|
320
|
+
stats.usage.output_tokens = (stats.usage.output_tokens || 0) + (usage.output_tokens || 0)
|
|
321
|
+
end
|
|
322
|
+
stats.usage.cost_usd = @total_cost
|
|
323
|
+
|
|
324
|
+
@result.stats = stats
|
|
325
|
+
end
|
|
326
|
+
|
|
327
|
+
#: () -> Array[String]
|
|
328
|
+
def command_line
|
|
329
|
+
command = if @base_command.is_a?(Array)
|
|
330
|
+
@base_command.dup
|
|
331
|
+
elsif @base_command.is_a?(String)
|
|
332
|
+
@base_command.split
|
|
333
|
+
else
|
|
334
|
+
["pi"]
|
|
335
|
+
end
|
|
336
|
+
command.push("--mode", "json", "-p")
|
|
337
|
+
command.push("--model", @model) if @model
|
|
338
|
+
command.push("--system-prompt", @replace_system_prompt) if @replace_system_prompt
|
|
339
|
+
command.push("--append-system-prompt", @append_system_prompt) if @append_system_prompt
|
|
340
|
+
if @session.present?
|
|
341
|
+
command.push("--fork", @session)
|
|
342
|
+
else
|
|
343
|
+
command.push("--no-session")
|
|
344
|
+
end
|
|
345
|
+
command
|
|
346
|
+
end
|
|
347
|
+
end
|
|
348
|
+
end
|
|
349
|
+
end
|
|
350
|
+
end
|
|
351
|
+
end
|
|
352
|
+
end
|
|
@@ -0,0 +1,41 @@
|
|
|
1
|
+
# typed: true
|
|
2
|
+
# frozen_string_literal: true
|
|
3
|
+
|
|
4
|
+
module Roast
|
|
5
|
+
module Cogs
|
|
6
|
+
class Agent < Cog
|
|
7
|
+
module Providers
|
|
8
|
+
class Pi < Provider
|
|
9
|
+
class Output < Agent::Output
|
|
10
|
+
delegate :response, :session, :stats, to: :@invocation_result
|
|
11
|
+
|
|
12
|
+
#: (PiInvocation::Result) -> void
|
|
13
|
+
def initialize(invocation_result)
|
|
14
|
+
super()
|
|
15
|
+
@invocation_result = invocation_result
|
|
16
|
+
end
|
|
17
|
+
end
|
|
18
|
+
|
|
19
|
+
#: (Agent::Input) -> Agent::Output
|
|
20
|
+
def invoke(input)
|
|
21
|
+
invocations = [] #: Array[PiInvocation]
|
|
22
|
+
input.prompts.each do |prompt|
|
|
23
|
+
previous_session = invocations.last&.result&.session
|
|
24
|
+
invocation = PiInvocation.new(
|
|
25
|
+
@config,
|
|
26
|
+
prompt,
|
|
27
|
+
previous_session || input.session,
|
|
28
|
+
)
|
|
29
|
+
invocation.run!
|
|
30
|
+
invocations << invocation
|
|
31
|
+
break unless invocation.result.success
|
|
32
|
+
end
|
|
33
|
+
final_result = invocations.last.not_nil!.result
|
|
34
|
+
final_result.stats = invocations.filter_map { |i| i.result.stats }.reduce(:+) if invocations.size > 1
|
|
35
|
+
Output.new(final_result)
|
|
36
|
+
end
|
|
37
|
+
end
|
|
38
|
+
end
|
|
39
|
+
end
|
|
40
|
+
end
|
|
41
|
+
end
|
|
@@ -61,6 +61,21 @@ module Roast
|
|
|
61
61
|
@model_usage = {}
|
|
62
62
|
end
|
|
63
63
|
|
|
64
|
+
# Add two Stats objects together, summing their durations, turns, usage, and model usage
|
|
65
|
+
#
|
|
66
|
+
# Nil values are treated as zero when the other operand is non-nil.
|
|
67
|
+
# Model usage hashes are merged, summing usage for models that appear in both.
|
|
68
|
+
#
|
|
69
|
+
#: (Stats) -> Stats
|
|
70
|
+
def +(other)
|
|
71
|
+
result = Stats.new
|
|
72
|
+
result.duration_ms = sum_nils(duration_ms, other.duration_ms)&.to_int
|
|
73
|
+
result.num_turns = sum_nils(num_turns, other.num_turns)&.to_int
|
|
74
|
+
result.usage = usage + other.usage
|
|
75
|
+
result.model_usage = merge_model_usage(model_usage, other.model_usage)
|
|
76
|
+
result
|
|
77
|
+
end
|
|
78
|
+
|
|
64
79
|
# Get a human-readable string representation of the statistics
|
|
65
80
|
#
|
|
66
81
|
# Formats the statistics into a multi-line string with the following information:
|
|
@@ -84,6 +99,20 @@ module Roast
|
|
|
84
99
|
end
|
|
85
100
|
lines.join("\n")
|
|
86
101
|
end
|
|
102
|
+
|
|
103
|
+
private
|
|
104
|
+
|
|
105
|
+
#: (Numeric?, Numeric?) -> Numeric?
|
|
106
|
+
def sum_nils(a, b)
|
|
107
|
+
return if a.nil? && b.nil?
|
|
108
|
+
|
|
109
|
+
(a || 0) + (b || 0)
|
|
110
|
+
end
|
|
111
|
+
|
|
112
|
+
#: (Hash[String, Usage], Hash[String, Usage]) -> Hash[String, Usage]
|
|
113
|
+
def merge_model_usage(a, b)
|
|
114
|
+
a.merge(b) { |_model, usage_a, usage_b| usage_a + usage_b }
|
|
115
|
+
end
|
|
87
116
|
end
|
|
88
117
|
end
|
|
89
118
|
end
|
|
@@ -54,6 +54,28 @@ module Roast
|
|
|
54
54
|
#
|
|
55
55
|
#: Float?
|
|
56
56
|
attr_accessor :cost_usd
|
|
57
|
+
|
|
58
|
+
# Add two Usage objects together, summing their token counts and costs
|
|
59
|
+
#
|
|
60
|
+
# Nil values are treated as zero when the other operand is non-nil.
|
|
61
|
+
#
|
|
62
|
+
#: (Usage) -> Usage
|
|
63
|
+
def +(other)
|
|
64
|
+
result = Usage.new
|
|
65
|
+
result.input_tokens = sum_nils(input_tokens, other.input_tokens)&.to_int
|
|
66
|
+
result.output_tokens = sum_nils(output_tokens, other.output_tokens)&.to_int
|
|
67
|
+
result.cost_usd = sum_nils(cost_usd, other.cost_usd)&.to_f
|
|
68
|
+
result
|
|
69
|
+
end
|
|
70
|
+
|
|
71
|
+
private
|
|
72
|
+
|
|
73
|
+
#: (Numeric?, Numeric?) -> Numeric?
|
|
74
|
+
def sum_nils(a, b)
|
|
75
|
+
return if a.nil? && b.nil?
|
|
76
|
+
|
|
77
|
+
(a || 0) + (b || 0)
|
|
78
|
+
end
|
|
57
79
|
end
|
|
58
80
|
end
|
|
59
81
|
end
|
data/lib/roast/cogs/agent.rb
CHANGED
|
@@ -47,13 +47,10 @@ module Roast
|
|
|
47
47
|
#
|
|
48
48
|
#: (Input) -> Output
|
|
49
49
|
def execute(input)
|
|
50
|
-
puts "[USER PROMPT] #{input.valid_prompt!}" if config.show_prompt?
|
|
51
50
|
output = provider.invoke(input)
|
|
52
|
-
|
|
53
|
-
|
|
54
|
-
|
|
55
|
-
puts "[AGENT STATS] #{output.stats}" if config.show_stats?
|
|
56
|
-
puts "Session ID: #{output.session}" if config.show_stats?
|
|
51
|
+
if config.show_stats?
|
|
52
|
+
Event << { block: { header: "AGENT STATS", content: "#{output.stats}\nSession ID: #{output.session}" } }
|
|
53
|
+
end
|
|
57
54
|
output
|
|
58
55
|
end
|
|
59
56
|
|
|
@@ -64,6 +61,8 @@ module Roast
|
|
|
64
61
|
@provider ||= case config.valid_provider!
|
|
65
62
|
when :claude
|
|
66
63
|
Providers::Claude.new(config)
|
|
64
|
+
when :pi
|
|
65
|
+
Providers::Pi.new(config)
|
|
67
66
|
else
|
|
68
67
|
raise UnknownProviderError, "Unknown provider: #{config.valid_provider!}"
|
|
69
68
|
end
|
|
@@ -12,6 +12,22 @@ module Roast
|
|
|
12
12
|
default_base_url: "https://api.openai.com/v1",
|
|
13
13
|
default_model: "gpt-4o-mini",
|
|
14
14
|
},
|
|
15
|
+
anthropic: {
|
|
16
|
+
api_key_env_var: "ANTHROPIC_API_KEY",
|
|
17
|
+
base_url_env_var: "ANTHROPIC_API_BASE",
|
|
18
|
+
default_base_url: "https://api.anthropic.com",
|
|
19
|
+
default_model: "claude-haiku-4-5",
|
|
20
|
+
},
|
|
21
|
+
perplexity: {
|
|
22
|
+
api_key_env_var: "PERPLEXITY_API_KEY",
|
|
23
|
+
default_model: "sonar",
|
|
24
|
+
},
|
|
25
|
+
gemini: {
|
|
26
|
+
api_key_env_var: "GEMINI_API_KEY",
|
|
27
|
+
base_url_env_var: "GEMINI_API_BASE",
|
|
28
|
+
default_base_url: "https://generativelanguage.googleapis.com/v1beta",
|
|
29
|
+
default_model: "gemini-3.1-flash-lite",
|
|
30
|
+
},
|
|
15
31
|
}.freeze #: Hash[Symbol, Hash[Symbol, String]]
|
|
16
32
|
|
|
17
33
|
# Configure the cog to use a specified API provider when invoking the llm
|
|
@@ -51,7 +67,7 @@ module Roast
|
|
|
51
67
|
def valid_provider!
|
|
52
68
|
provider = @values[:provider] || PROVIDERS.keys.first
|
|
53
69
|
unless PROVIDERS.include?(provider)
|
|
54
|
-
raise
|
|
70
|
+
raise InvalidConfigError, "'#{provider}' is not a valid provider. Available providers include: #{PROVIDERS.keys.join(", ")}"
|
|
55
71
|
end
|
|
56
72
|
|
|
57
73
|
provider
|
|
@@ -75,6 +91,9 @@ module Roast
|
|
|
75
91
|
#
|
|
76
92
|
# #### Environment Variables
|
|
77
93
|
# - OpenAI Provider: OPENAI_API_KEY
|
|
94
|
+
# - Anthropic Provider: ANTHROPIC_API_KEY
|
|
95
|
+
# - Perplexity Provider: PERPLEXITY_API_KEY
|
|
96
|
+
# - Gemini Provider: GEMINI_API_KEY
|
|
78
97
|
#
|
|
79
98
|
# #### See Also
|
|
80
99
|
# - `api_key`
|
|
@@ -91,6 +110,9 @@ module Roast
|
|
|
91
110
|
#
|
|
92
111
|
# #### Environment Variables
|
|
93
112
|
# - OpenAI Provider: OPENAI_API_KEY
|
|
113
|
+
# - Anthropic Provider: ANTHROPIC_API_KEY
|
|
114
|
+
# - Perplexity Provider: PERPLEXITY_API_KEY
|
|
115
|
+
# - Gemini Provider: GEMINI_API_KEY
|
|
94
116
|
#
|
|
95
117
|
# #### See Also
|
|
96
118
|
# - `api_key`
|
|
@@ -126,6 +148,8 @@ module Roast
|
|
|
126
148
|
#
|
|
127
149
|
# #### Environment Variables
|
|
128
150
|
# - OpenAI Provider: OPENAI_API_BASE
|
|
151
|
+
# - Anthropic Provider: ANTHROPIC_API_BASE
|
|
152
|
+
# - Gemini Provider: GEMINI_API_BASE
|
|
129
153
|
#
|
|
130
154
|
# #### See Also
|
|
131
155
|
# - `base_url`
|
|
@@ -139,6 +163,8 @@ module Roast
|
|
|
139
163
|
#
|
|
140
164
|
# #### Environment Variables
|
|
141
165
|
# - OpenAI Provider: OPENAI_API_BASE
|
|
166
|
+
# - Anthropic Provider: ANTHROPIC_API_BASE
|
|
167
|
+
# - Gemini Provider: GEMINI_API_BASE
|
|
142
168
|
#
|
|
143
169
|
# #### See Also
|
|
144
170
|
# - `base_url`
|
|
@@ -171,7 +197,7 @@ module Roast
|
|
|
171
197
|
#
|
|
172
198
|
#: () -> void
|
|
173
199
|
def use_default_model!
|
|
174
|
-
@values
|
|
200
|
+
@values.delete(:model)
|
|
175
201
|
end
|
|
176
202
|
|
|
177
203
|
# Get the validated, configured value of the model the cog is configured to use when running the agent
|