Skip to content

Commit ef24b1b

Browse files
Merge pull request #36 from openlayer-ai/vini/open-11569-support-additional-columns-in-google-conversational-search
Support additional columns in GoogleConversationalSearchTracer
2 parents 0aebb3a + e5cd8df commit ef24b1b

4 files changed

Lines changed: 233 additions & 11 deletions

File tree

examples/google_tracer.rb

Lines changed: 10 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
# Add lib directory to load path
88
$LOAD_PATH.unshift(File.expand_path("../lib", __dir__))
99

10+
require "securerandom"
1011
require "openlayer"
1112
require "openlayer/integrations/google_conversational_search_tracer"
1213
require "google/cloud/discovery_engine/v1"
@@ -19,19 +20,25 @@
1920
api_key: ENV["OPENLAYER_API_KEY"]
2021
)
2122

22-
# Enable tracing - this patches the client to send all queries to Openlayer
23+
# Enable tracing - this patches the client to send all queries to Openlayer.
24+
# additional_columns here is a static default applied to every trace sent
25+
# through this client.
2326
Openlayer::Integrations::GoogleConversationalSearchTracer.trace_client(
2427
google_client,
2528
openlayer_client: openlayer,
26-
inference_pipeline_id: ENV["OPENLAYER_INFERENCE_PIPELINE_ID"]
29+
inference_pipeline_id: ENV["OPENLAYER_INFERENCE_PIPELINE_ID"],
30+
additional_columns: {environment: "production"}
2731
)
2832

2933
# Use the client normally - all answer_query calls are now automatically traced!
34+
# additional_columns here is per-call; it takes precedence over the static
35+
# default above on a key conflict.
3036
response = google_client.answer_query(
3137
serving_config: ENV["GOOGLE_SERVING_CONFIG"],
3238
query: Google::Cloud::DiscoveryEngine::V1::Query.new(
3339
text: "What is the meaning of life?"
34-
)
40+
),
41+
additional_columns: {trace_id: SecureRandom.uuid}
3542
)
3643

3744
puts "Answer: #{response.answer.answer_text}"

lib/openlayer/integrations/google_conversational_search_tracer.rb

Lines changed: 75 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -21,15 +21,25 @@ module Integrations
2121
# Openlayer::Integrations::GoogleConversationalSearchTracer.trace_client(
2222
# google_client,
2323
# openlayer_client: openlayer,
24-
# inference_pipeline_id: 'your-pipeline-id'
24+
# inference_pipeline_id: 'your-pipeline-id',
25+
# additional_columns: { environment: 'production' }
2526
# )
2627
#
27-
# # Now all answer_query calls are automatically traced
28+
# # Now all answer_query calls are automatically traced! Pass
29+
# # additional_columns on an individual call to attach data (like your
30+
# # own trace ID) to just that row; it takes precedence over the
31+
# # static defaults above on a key conflict.
2832
# response = google_client.answer_query(
2933
# serving_config: "projects/.../servingConfigs/default",
30-
# query: { text: "What is the meaning of life?" }
34+
# query: { text: "What is the meaning of life?" },
35+
# additional_columns: { trace_id: "abc-123" }
3136
# )
3237
class GoogleConversationalSearchTracer
38+
# Row keys computed by this tracer. Any key in a caller-supplied
39+
# additional_columns hash matching one of these is dropped, so custom
40+
# data can never overwrite core trace fields.
41+
RESERVED_ROW_KEYS = [:query, :answer, :latency_ms, :timestamp, :metadata, :steps, :context, :session_id, :user_id].freeze
42+
3343
# Enable tracing on a Google ConversationalSearchService client
3444
#
3545
# @param client [Google::Cloud::DiscoveryEngine::V1::ConversationalSearchService::Client]
@@ -42,8 +52,12 @@ class GoogleConversationalSearchTracer
4252
# Optional session ID to use for all traces. Takes precedence over auto-extracted sessions.
4353
# @param user_id [String, nil]
4454
# Optional user ID to use for all traces.
55+
# @param additional_columns [Hash, nil]
56+
# Optional static column values merged into every trace sent through this client (e.g. `{ environment: 'production' }`).
57+
# A value passed to an individual answer_query call takes precedence over these on a key conflict. Keys colliding
58+
# with a reserved row column (query, answer, latency_ms, timestamp, metadata, steps, context, session_id, user_id) are dropped.
4559
# @return [void]
46-
def self.trace_client(client, openlayer_client:, inference_pipeline_id:, session_id: nil, user_id: nil)
60+
def self.trace_client(client, openlayer_client:, inference_pipeline_id:, session_id: nil, user_id: nil, additional_columns: {})
4761
# Store original method reference
4862
original_answer_query = client.method(:answer_query)
4963

@@ -52,6 +66,10 @@ def self.trace_client(client, openlayer_client:, inference_pipeline_id:, session
5266
# Capture start time
5367
start_time = Time.now
5468

69+
# Extract per-call additional columns before forwarding to the
70+
# real client; Google's client never sees this key
71+
call_additional_columns = kwargs.delete(:additional_columns)
72+
5573
# Execute the original method
5674
response = original_answer_query.call(*args, **kwargs, &block)
5775

@@ -69,7 +87,9 @@ def self.trace_client(client, openlayer_client:, inference_pipeline_id:, session
6987
openlayer_client: openlayer_client,
7088
inference_pipeline_id: inference_pipeline_id,
7189
session_id: session_id,
72-
user_id: user_id
90+
user_id: user_id,
91+
additional_columns: additional_columns,
92+
call_additional_columns: call_additional_columns
7393
)
7494
rescue StandardError => e
7595
# Never break the user's application due to tracing errors
@@ -95,8 +115,10 @@ def self.trace_client(client, openlayer_client:, inference_pipeline_id:, session
95115
# @param inference_pipeline_id [String] Pipeline ID
96116
# @param session_id [String, nil] Optional session ID (takes precedence over auto-extracted)
97117
# @param user_id [String, nil] Optional user ID
118+
# @param additional_columns [Hash, nil] Optional static column values (see {.trace_client})
119+
# @param call_additional_columns [Hash, nil] Optional per-call column values; takes precedence over additional_columns
98120
# @return [void]
99-
def self.send_trace(args:, kwargs:, response:, start_time:, end_time:, openlayer_client:, inference_pipeline_id:, session_id: nil, user_id: nil)
121+
def self.send_trace(args:, kwargs:, response:, start_time:, end_time:, openlayer_client:, inference_pipeline_id:, session_id: nil, user_id: nil, additional_columns: {}, call_additional_columns: {})
100122
# Calculate latency
101123
latency_ms = ((end_time - start_time) * 1000).round(2)
102124

@@ -199,6 +221,12 @@ def self.send_trace(args:, kwargs:, response:, start_time:, end_time:, openlayer
199221
trace_data[:config][:userIdColumnName] = "user_id"
200222
end
201223

224+
# Merge additional columns (per-call values take precedence over
225+
# static defaults; keys colliding with reserved row columns are
226+
# dropped so custom data can never corrupt core trace fields)
227+
extra_columns = resolve_additional_columns(additional_columns, call_additional_columns)
228+
trace_data[:rows][0].merge!(extra_columns) unless extra_columns.empty?
229+
202230
# Send to Openlayer
203231
openlayer_client
204232
.inference_pipelines
@@ -594,6 +622,45 @@ def self.extract_query_understanding_info(answer)
594622
nil
595623
end
596624

625+
# Merge static and per-call additional columns into a single Hash of
626+
# extra row columns. Call-level values take precedence over static
627+
# ones on key conflict, and any key colliding with a reserved row
628+
# column is dropped.
629+
#
630+
# @param static_columns [Object] Value passed to trace_client (expected Hash)
631+
# @param call_columns [Object] Value passed to an individual answer_query call (expected Hash)
632+
# @return [Hash] Extra columns safe to merge onto a trace row
633+
def self.resolve_additional_columns(static_columns, call_columns)
634+
merged = normalize_additional_columns(static_columns).merge(normalize_additional_columns(call_columns))
635+
636+
merged.each_with_object({}) do |(key, value), result|
637+
if RESERVED_ROW_KEYS.include?(key)
638+
warn_if_debug("[Openlayer] additional_columns key :#{key} collides with a reserved column and was ignored")
639+
else
640+
result[key] = value
641+
end
642+
end
643+
end
644+
645+
# Normalize an additional_columns value into a Hash with Symbol keys.
646+
# Non-Hash input (or a key that can't be a Symbol) is dropped rather
647+
# than raising, so a caller mistake can never break tracing.
648+
#
649+
# @param columns [Object] Expected to be a Hash of column name => value
650+
# @return [Hash]
651+
def self.normalize_additional_columns(columns)
652+
return {} unless columns.is_a?(Hash)
653+
654+
columns.each_with_object({}) do |(key, value), result|
655+
next unless key.respond_to?(:to_sym)
656+
657+
result[key.to_sym] = value
658+
end
659+
rescue StandardError => e
660+
warn_if_debug("[Openlayer] Failed to normalize additional columns: #{e.message}")
661+
{}
662+
end
663+
597664
# Safely extract a field from an object
598665
#
599666
# @param obj [Object] Object to extract from
@@ -659,6 +726,8 @@ def self.warn_if_debug(message)
659726
:extract_session,
660727
:extract_user_pseudo_id,
661728
:extract_query_understanding_info,
729+
:resolve_additional_columns,
730+
:normalize_additional_columns,
662731
:safe_extract,
663732
:safe_count,
664733
:extract_timestamp

rbi/openlayer/integrations.rbi

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -10,15 +10,17 @@ module Openlayer
1010
openlayer_client: Openlayer::Client,
1111
inference_pipeline_id: String,
1212
session_id: T.nilable(String),
13-
user_id: T.nilable(String)
13+
user_id: T.nilable(String),
14+
additional_columns: T::Hash[Symbol, T.untyped]
1415
).void
1516
end
1617
def self.trace_client(
1718
client,
1819
openlayer_client:,
1920
inference_pipeline_id:,
2021
session_id: nil,
21-
user_id: nil
22+
user_id: nil,
23+
additional_columns: {}
2224
)
2325
end
2426
end
Lines changed: 144 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,144 @@
1+
# frozen_string_literal: true
2+
3+
require_relative "../test_helper"
4+
require_relative "../../../lib/openlayer/integrations/google_conversational_search_tracer"
5+
6+
module Openlayer
7+
module Test
8+
module Integrations
9+
end
10+
end
11+
end
12+
13+
class Openlayer::Test::Integrations::GoogleConversationalSearchTracerTest < Minitest::Test
14+
Tracer = Openlayer::Integrations::GoogleConversationalSearchTracer
15+
16+
class FakeAnswer
17+
attr_reader :answer_text
18+
19+
def initialize(answer_text)
20+
@answer_text = answer_text
21+
end
22+
end
23+
24+
class FakeResponse
25+
attr_reader :answer
26+
27+
def initialize(answer_text)
28+
@answer = FakeAnswer.new(answer_text)
29+
end
30+
end
31+
32+
class FakeGoogleClient
33+
def answer_query(serving_config:, query:) # rubocop:disable Lint/UnusedMethodArgument
34+
FakeResponse.new("hi")
35+
end
36+
end
37+
38+
class FakeDataResource
39+
attr_reader :calls
40+
41+
def initialize
42+
@calls = []
43+
end
44+
45+
def stream(inference_pipeline_id, **trace_data)
46+
@calls << {inference_pipeline_id: inference_pipeline_id}.merge(trace_data)
47+
end
48+
end
49+
50+
class FakeInferencePipelines
51+
attr_reader :data
52+
53+
def initialize(data)
54+
@data = data
55+
end
56+
end
57+
58+
class FakeOpenlayerClient
59+
attr_reader :inference_pipelines
60+
61+
def initialize
62+
@data = FakeDataResource.new
63+
@inference_pipelines = FakeInferencePipelines.new(@data)
64+
end
65+
66+
def last_row
67+
@data.calls.last[:rows][0]
68+
end
69+
end
70+
71+
def setup
72+
@openlayer_client = FakeOpenlayerClient.new
73+
@start_time = Time.now
74+
@end_time = @start_time + 1
75+
end
76+
77+
def trace_row(**overrides)
78+
defaults = {
79+
args: [],
80+
kwargs: {query: "hello"},
81+
response: FakeResponse.new("hi"),
82+
start_time: @start_time,
83+
end_time: @end_time,
84+
openlayer_client: @openlayer_client,
85+
inference_pipeline_id: "pipeline-id"
86+
}
87+
88+
Tracer.send_trace(**defaults, **overrides)
89+
@openlayer_client.last_row
90+
end
91+
92+
def test_static_additional_columns_appear_on_the_row
93+
row = trace_row(additional_columns: {environment: "production"})
94+
95+
assert_equal("production", row[:environment])
96+
end
97+
98+
def test_per_call_additional_columns_appear_on_the_row
99+
row = trace_row(call_additional_columns: {trace_id: "abc-123"})
100+
101+
assert_equal("abc-123", row[:trace_id])
102+
end
103+
104+
def test_per_call_additional_columns_override_static_on_conflict
105+
row = trace_row(
106+
additional_columns: {trace_id: "static-value"},
107+
call_additional_columns: {trace_id: "call-value"}
108+
)
109+
110+
assert_equal("call-value", row[:trace_id])
111+
end
112+
113+
def test_reserved_keys_are_dropped_even_as_string_keys
114+
row = trace_row(additional_columns: {"answer" => "hijacked", trace_id: "abc-123"})
115+
116+
assert_equal("hi", row[:answer])
117+
assert_equal("abc-123", row[:trace_id])
118+
end
119+
120+
def test_non_hash_additional_columns_does_not_raise
121+
row = trace_row(additional_columns: "not-a-hash", call_additional_columns: nil)
122+
123+
assert_equal("hi", row[:answer])
124+
end
125+
126+
def test_trace_client_strips_additional_columns_before_forwarding_to_google_client
127+
google_client = FakeGoogleClient.new
128+
129+
Tracer.trace_client(
130+
google_client,
131+
openlayer_client: @openlayer_client,
132+
inference_pipeline_id: "pipeline-id"
133+
)
134+
135+
response = google_client.answer_query(
136+
serving_config: "config",
137+
query: "hello",
138+
additional_columns: {trace_id: "abc-123"}
139+
)
140+
141+
assert_equal("hi", response.answer.answer_text)
142+
assert_equal("abc-123", @openlayer_client.last_row[:trace_id])
143+
end
144+
end

0 commit comments

Comments
 (0)