diff --git a/lib/ai.rb b/lib/ai.rb index 5bbdf46..a00578d 100644 --- a/lib/ai.rb +++ b/lib/ai.rb @@ -47,7 +47,7 @@ class Error < StandardError LogProbs = T.type_alias { T.anything } ProviderMetadata = T.type_alias { T.anything } - config_accessor :origin, :client, :api_key + config_accessor :origin, :client, :api_key, :delegated_token_resolver sig { params(content: String).returns(Ai::Message) } def self.user_message(content) diff --git a/lib/ai/agent.rb b/lib/ai/agent.rb index 70274c4..81e93f7 100644 --- a/lib/ai/agent.rb +++ b/lib/ai/agent.rb @@ -22,7 +22,8 @@ def initialize(agent_name:, client: Ai.client) runtime_context: T::Hash[String, T.anything], max_retries: Integer, max_steps: Integer, - telemetry: Ai::TelemetrySettings + telemetry: Ai::TelemetrySettings, + delegated_auth: T.nilable(Object) ).returns(Ai::GenerateTextResult) end def generate_text( @@ -30,7 +31,8 @@ def generate_text( runtime_context: {}, max_retries: 2, max_steps: 5, - telemetry: Ai::TelemetrySettings.new + telemetry: Ai::TelemetrySettings.new, + delegated_auth: nil ) options = { runtime_context: runtime_context, @@ -39,7 +41,8 @@ def generate_text( telemetry: telemetry } - data = client.generate(agent_name, messages: messages, options: options) + delegated_token = resolve_delegated_token(delegated_auth) + data = client.generate(agent_name, messages: messages, options: options, delegated_token: delegated_token) TypeCoerce[Ai::GenerateTextResult].new.from(data, raise_coercion_error: false) end @@ -51,7 +54,8 @@ def generate_text( runtime_context: T::Hash[String, T.anything], max_retries: Integer, max_steps: Integer, - telemetry: Ai::TelemetrySettings + telemetry: Ai::TelemetrySettings, + delegated_auth: T.nilable(Object) ) .returns(GenerateObjectResult[T.type_parameter(:O)]) end @@ -61,7 +65,8 @@ def generate_object( runtime_context: {}, max_retries: 2, max_steps: 5, - telemetry: Ai::TelemetrySettings.new + telemetry: Ai::TelemetrySettings.new, + delegated_auth: nil ) schema = Ai::StructToJsonSchema.convert(T.cast(output_class, T.class_of(T::Struct))) @@ -75,7 +80,8 @@ def generate_object( telemetry: telemetry } - data = client.generate(agent_name, messages: messages, options: options) + delegated_token = resolve_delegated_token(delegated_auth) + data = client.generate(agent_name, messages: messages, options: options, delegated_token: delegated_token) object = TypeCoerce[output_class].from(data['object']) TypeCoerce[GenerateObjectResult] @@ -83,5 +89,17 @@ def generate_object( .from(data, raise_coercion_error: false) .with(object: object) end + + private + + sig { params(delegated_auth: T.nilable(Object)).returns(T.nilable(String)) } + def resolve_delegated_token(delegated_auth) + return nil if delegated_auth.nil? + + resolver = Ai.config.delegated_token_resolver + raise Ai::Error, 'delegated_token_resolver is not configured. Set Ai.delegated_token_resolver in your initializer.' if resolver.nil? + + resolver.call(delegated_auth) + end end end diff --git a/lib/ai/client.rb b/lib/ai/client.rb index ffe2631..36bd209 100644 --- a/lib/ai/client.rb +++ b/lib/ai/client.rb @@ -32,15 +32,16 @@ def agent_names .params( agent_name: String, messages: T::Array[Ai::Message], - options: T::Hash[Symbol, T.anything] + options: T::Hash[Symbol, T.anything], + delegated_token: T.nilable(String) ) .returns(T::Hash[String, T.anything]) end - def generate(agent_name, messages:, options: {}) + def generate(agent_name, messages:, options: {}, delegated_token: nil) end - sig { abstract.params(workflow_name: String, input: T::Struct).returns(ApiResponse) } - def run_workflow(workflow_name, input:) + sig { abstract.params(workflow_name: String, input: T::Struct, delegated_token: T.nilable(String)).returns(ApiResponse) } + def run_workflow(workflow_name, input:, delegated_token: nil) end sig { abstract.params(workflow_name: String).returns(SchemaHash) } diff --git a/lib/ai/clients/mastra.rb b/lib/ai/clients/mastra.rb index c19b62d..ecfcb58 100644 --- a/lib/ai/clients/mastra.rb +++ b/lib/ai/clients/mastra.rb @@ -60,13 +60,14 @@ def agent_names .params( agent_name: String, messages: T::Array[Ai::Message], - options: T::Hash[Symbol, T.anything] + options: T::Hash[Symbol, T.anything], + delegated_token: T.nilable(String) ) .returns(T::Hash[String, T.anything]) end - def generate(agent_name, messages:, options: {}) + def generate(agent_name, messages:, options: {}, delegated_token: nil) url = URI.join(@base_uri, "api/agents/#{agent_name}/generate") - generated_response = response(url: url, messages: messages, options: options) + generated_response = response(url: url, messages: messages, options: options, delegated_token: delegated_token) parsed_response = JSON.parse(generated_response.body || '').deep_transform_keys(&:underscore) @@ -83,15 +84,15 @@ def generate(agent_name, messages:, options: {}) end sig do - override.params(workflow_name: String, input: T::Struct).returns(Ai::Client::ApiResponse) + override.params(workflow_name: String, input: T::Struct, delegated_token: T.nilable(String)).returns(Ai::Client::ApiResponse) end - def run_workflow(workflow_name, input:) + def run_workflow(workflow_name, input:, delegated_token: nil) run_id = SecureRandom.uuid # Step 1: Create a new run for the workflow create_url = URI.join(@base_uri, "api/workflows/#{workflow_name}/create-run?runId=#{run_id}") - create_response = http_post(create_url, body: '{}') + create_response = http_post(create_url, body: '{}', delegated_token: delegated_token) unless create_response.is_a?(Net::HTTPSuccess) raise Ai::Error, "Mastra error – could not create workflow run: #{create_response.body}" @@ -104,7 +105,7 @@ def run_workflow(workflow_name, input:) stream_error_body = T.let(nil, T.nilable(String)) stream_body_chunks = T.let([], T::Array[String]) stream_response = - http_post(stream_url, body: stream_request_body, stream: true) do |response| + http_post(stream_url, body: stream_request_body, stream: true, delegated_token: delegated_token) do |response| if response.is_a?(Net::HTTPSuccess) response.read_body do |chunk| # Capture the stream body to check for workflow failures @@ -144,6 +145,7 @@ def run_workflow(workflow_name, input:) .config .api_key .present? + result_request['X-Factorial-Delegated-Bearer'] = delegated_token if delegated_token.present? http = build_http result_response = http.request(result_request) @@ -244,13 +246,15 @@ def deep_camelize_keys(options) url: URI::Generic, body: T.nilable(String), stream: T::Boolean, + delegated_token: T.nilable(String), blk: T.nilable(T.proc.params(response: Net::HTTPResponse).void) ).returns(Net::HTTPResponse) end - def http_post(url, body: nil, stream: false, &blk) + def http_post(url, body: nil, stream: false, delegated_token: nil, &blk) request = Net::HTTP::Post.new(url) request['Origin'] = Ai.config.origin request['Authorization'] = "Bearer #{Ai.config.api_key}" if Ai.config.api_key.present? + request['X-Factorial-Delegated-Bearer'] = delegated_token if delegated_token.present? if body request['Content-Type'] = 'application/json' request.body = body @@ -282,14 +286,16 @@ def http_post(url, body: nil, stream: false, &blk) params( url: URI::Generic, messages: T::Array[Ai::Message], - options: T::Hash[Symbol, T.anything] + options: T::Hash[Symbol, T.anything], + delegated_token: T.nilable(String) ).returns(Net::HTTPResponse) end - def response(url:, messages:, options:) + def response(url:, messages:, options:, delegated_token: nil) request = Net::HTTP::Post.new(url) request['Content-Type'] = 'application/json' request['Origin'] = Ai.config.origin request['Authorization'] = "Bearer #{Ai.config.api_key}" if Ai.config.api_key.present? + request['X-Factorial-Delegated-Bearer'] = delegated_token if delegated_token.present? # convert to camelCase and unpacking for API compatibility camelized_options = deep_camelize_keys(options) diff --git a/lib/ai/clients/test.rb b/lib/ai/clients/test.rb index 52203d7..a98c693 100644 --- a/lib/ai/clients/test.rb +++ b/lib/ai/clients/test.rb @@ -30,11 +30,12 @@ def set_returned_object(output) # rubocop:disable Naming/AccessorMethodName .params( agent_name: String, messages: T::Array[Ai::Message], - options: T::Hash[Symbol, T.anything] + options: T::Hash[Symbol, T.anything], + delegated_token: T.nilable(String) ) .returns(T::Hash[String, T.anything]) end - def generate(agent_name, messages:, options: {}) + def generate(agent_name, messages:, options: {}, delegated_token: nil) output = options[:structured_output] # Use the first message content for testing purposes @@ -101,9 +102,9 @@ def generate(agent_name, messages:, options: {}) end sig do - override.params(workflow_name: String, input: T::Struct).returns(Ai::Client::ApiResponse) + override.params(workflow_name: String, input: T::Struct, delegated_token: T.nilable(String)).returns(Ai::Client::ApiResponse) end - def run_workflow(workflow_name, input:) + def run_workflow(workflow_name, input:, delegated_token: nil) @returned_object end diff --git a/spec/lib/ai/agent_spec.rb b/spec/lib/ai/agent_spec.rb index ea68830..dfecff1 100644 --- a/spec/lib/ai/agent_spec.rb +++ b/spec/lib/ai/agent_spec.rb @@ -28,7 +28,8 @@ max_retries: 2, max_steps: 5, telemetry: anything - } + }, + delegated_token: nil ).and_call_original agent.generate_text(messages: [Ai.user_message('Hello')], runtime_context: runtime_context) @@ -44,7 +45,8 @@ max_retries: 5, max_steps: 10, telemetry: anything - } + }, + delegated_token: nil ).and_call_original agent.generate_text(messages: [Ai.user_message('Hello')], max_retries: 5, max_steps: 10) @@ -71,7 +73,8 @@ max_retries: 2, max_steps: 5, telemetry: telemetry_settings - } + }, + delegated_token: nil ).and_call_original agent.generate_text(messages: [Ai.user_message('Hello')], telemetry: telemetry_settings) @@ -87,7 +90,8 @@ max_retries: 2, max_steps: 5, telemetry: kind_of(Ai::TelemetrySettings) - } + }, + delegated_token: nil ).and_call_original result = agent.generate_text(messages: [Ai.user_message('Hello')]) @@ -111,6 +115,45 @@ expect(result.tool_results).to eq([]) expect(result.steps).to eq([]) end + + describe 'delegated_auth' do + after { Ai.delegated_token_resolver = nil } + + it 'resolves delegated_auth to a token via the configured resolver and passes it to the client' do + principal = instance_double(Object) + Ai.delegated_token_resolver = ->(p) { "resolved-token-for-#{p.class.name}" } + + expect(client).to receive(:generate).with( + 'test', + messages: anything, + options: anything, + delegated_token: "resolved-token-for-#{principal.class.name}" + ).and_call_original + + agent.generate_text(messages: [Ai.user_message('Hello')], delegated_auth: principal) + end + + it 'passes nil delegated_token when delegated_auth is nil' do + Ai.delegated_token_resolver = ->(_p) { raise 'should not be called' } + + expect(client).to receive(:generate).with( + 'test', + messages: anything, + options: anything, + delegated_token: nil + ).and_call_original + + agent.generate_text(messages: [Ai.user_message('Hello')]) + end + + it 'raises Ai::Error when resolver is not configured but delegated_auth is provided' do + Ai.delegated_token_resolver = nil + + expect do + agent.generate_text(messages: [Ai.user_message('Hello')], delegated_auth: instance_double(Object)) + end.to raise_error(Ai::Error, /delegated_token_resolver is not configured/) + end + end end describe '#generate_object' do @@ -153,7 +196,8 @@ max_steps: 8, structured_output: hash_including(schema: anything), telemetry: anything - ) + ), + delegated_token: nil ).and_call_original agent.generate_object( @@ -184,7 +228,8 @@ hash_including( telemetry: telemetry_settings, structured_output: hash_including(schema: anything) - ) + ), + delegated_token: nil ).and_call_original agent.generate_object( @@ -217,5 +262,36 @@ expect { subject }.to raise_error(ArgumentError) end end + + describe 'delegated_auth' do + after { Ai.delegated_token_resolver = nil } + + before { client.set_returned_object({ 'name' => 'John Doe', 'age' => 30 }) } + + let(:schema) do + Class.new(T::Struct) do + const :name, String + const :age, Integer + end + end + + it 'resolves delegated_auth and passes delegated_token to the client' do + principal = instance_double(Object) + Ai.delegated_token_resolver = ->(_p) { 'company-token-abc' } + + expect(client).to receive(:generate).with( + 'test', + messages: anything, + options: anything, + delegated_token: 'company-token-abc' + ).and_call_original + + agent.generate_object( + messages: [Ai.user_message('Create person')], + output_class: schema, + delegated_auth: principal + ) + end + end end end diff --git a/spec/lib/ai/mastra_client_spec.rb b/spec/lib/ai/mastra_client_spec.rb index 0b0b73d..f911978 100644 --- a/spec/lib/ai/mastra_client_spec.rb +++ b/spec/lib/ai/mastra_client_spec.rb @@ -123,6 +123,51 @@ expect(stub).to have_been_requested end + + it 'sends X-Factorial-Delegated-Bearer header when delegated_token is provided' do + delegated_token = 'delegated-jwt-token-abc123' + + stub = + stub_request(:post, 'https://mastra.local.factorial.dev/api/agents/marvin/generate') + .with do |req| + req.headers['X-Factorial-Delegated-Bearer'] == delegated_token + end + .to_return( + status: 200, + body: { text: 'Test response' }.to_json, + headers: { 'Content-Type' => 'application/json' } + ) + + client.generate( + 'marvin', + messages: [Ai.user_message('test')], + options: {}, + delegated_token: delegated_token + ) + + expect(stub).to have_been_requested + end + + it 'does not send X-Factorial-Delegated-Bearer header when delegated_token is not provided' do + stub = + stub_request(:post, 'https://mastra.local.factorial.dev/api/agents/marvin/generate') + .with do |req| + !req.headers.key?('X-Factorial-Delegated-Bearer') + end + .to_return( + status: 200, + body: { text: 'Test response' }.to_json, + headers: { 'Content-Type' => 'application/json' } + ) + + client.generate( + 'marvin', + messages: [Ai.user_message('test')], + options: {} + ) + + expect(stub).to have_been_requested + end end describe '#run_workflow' do @@ -144,5 +189,39 @@ expect(result).to eq('sumOfNumbers' => 8) end end + + it 'sends X-Factorial-Delegated-Bearer header on all requests when delegated_token is provided' do + delegated_token = 'delegated-jwt-token-workflow' + + input = Class.new(T::Struct) { const :value, Integer }.new(value: 1) + + # Stub create-run + create_stub = + stub_request(:post, %r{mastra\.local\.factorial\.dev/api/workflows/#{workflow_name}/create-run}) + .with { |req| req.headers['X-Factorial-Delegated-Bearer'] == delegated_token } + .to_return(status: 200, body: '{}') + + # Stub stream + stream_stub = + stub_request(:post, %r{mastra\.local\.factorial\.dev/api/workflows/#{workflow_name}/stream}) + .with { |req| req.headers['X-Factorial-Delegated-Bearer'] == delegated_token } + .to_return(status: 200, body: '{"workflowStatus":"success"}') + + # Stub result fetch + result_stub = + stub_request(:get, %r{mastra\.local\.factorial\.dev/api/workflows/#{workflow_name}/runs/}) + .with { |req| req.headers['X-Factorial-Delegated-Bearer'] == delegated_token } + .to_return( + status: 200, + body: { status: 'success', result: { output: 'done' } }.to_json, + headers: { 'Content-Type' => 'application/json' } + ) + + client.run_workflow(workflow_name, input: input, delegated_token: delegated_token) + + expect(create_stub).to have_been_requested + expect(stream_stub).to have_been_requested + expect(result_stub).to have_been_requested + end end end