revenium-python-sdk 0.1.0__py3-none-any.whl
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.
- revenium_middleware/__init__.py +184 -0
- revenium_middleware/_core/__init__.py +65 -0
- revenium_middleware/_core/config.py +165 -0
- revenium_middleware/_core/context.py +109 -0
- revenium_middleware/_core/decorators.py +202 -0
- revenium_middleware/_core/metering.py +207 -0
- revenium_middleware/_core/prompt_extraction.py +55 -0
- revenium_middleware/_core/subscriber.py +51 -0
- revenium_middleware/_core/trace_fields.py +265 -0
- revenium_middleware/anthropic/__init__.py +108 -0
- revenium_middleware/anthropic/bedrock_adapter.py +753 -0
- revenium_middleware/anthropic/config.py +29 -0
- revenium_middleware/anthropic/middleware.py +1070 -0
- revenium_middleware/anthropic/prompt_extractor.py +178 -0
- revenium_middleware/anthropic/provider.py +141 -0
- revenium_middleware/anthropic/summary_printer.py +286 -0
- revenium_middleware/anthropic/trace_fields.py +158 -0
- revenium_middleware/google/__init__.py +114 -0
- revenium_middleware/google/common/__init__.py +127 -0
- revenium_middleware/google/common/exceptions.py +137 -0
- revenium_middleware/google/common/protocols.py +192 -0
- revenium_middleware/google/common/summary_printer.py +271 -0
- revenium_middleware/google/common/trace_fields.py +205 -0
- revenium_middleware/google/common/types.py +208 -0
- revenium_middleware/google/common/utils.py +1111 -0
- revenium_middleware/google/config.py +64 -0
- revenium_middleware/google/google_ai/__init__.py +53 -0
- revenium_middleware/google/google_ai/middleware.py +667 -0
- revenium_middleware/google/google_ai/provider.py +135 -0
- revenium_middleware/google/prompt_extractor.py +396 -0
- revenium_middleware/google/vertex_ai/__init__.py +56 -0
- revenium_middleware/google/vertex_ai/middleware.py +1162 -0
- revenium_middleware/google/vertex_ai/provider.py +99 -0
- revenium_middleware/litellm/__init__.py +25 -0
- revenium_middleware/litellm/client/__init__.py +81 -0
- revenium_middleware/litellm/client/config.py +53 -0
- revenium_middleware/litellm/client/context.py +198 -0
- revenium_middleware/litellm/client/decorators.py +912 -0
- revenium_middleware/litellm/client/hooks.py +192 -0
- revenium_middleware/litellm/client/integrations/__init__.py +26 -0
- revenium_middleware/litellm/client/integrations/crewai.py +446 -0
- revenium_middleware/litellm/client/middleware.py +321 -0
- revenium_middleware/litellm/client/summary_printer.py +314 -0
- revenium_middleware/litellm/client/trace_fields.py +51 -0
- revenium_middleware/litellm/client/validation.py +207 -0
- revenium_middleware/litellm/proxy/__init__.py +25 -0
- revenium_middleware/litellm/proxy/middleware.py +217 -0
- revenium_middleware/ollama/__init__.py +28 -0
- revenium_middleware/ollama/middleware.py +569 -0
- revenium_middleware/ollama/trace_fields.py +63 -0
- revenium_middleware/openai/__init__.py +23 -0
- revenium_middleware/openai/azure_config.py +169 -0
- revenium_middleware/openai/azure_model_resolver.py +219 -0
- revenium_middleware/openai/config.py +45 -0
- revenium_middleware/openai/exceptions.py +115 -0
- revenium_middleware/openai/langchain/__init__.py +114 -0
- revenium_middleware/openai/langchain/_utils.py +129 -0
- revenium_middleware/openai/langchain/unified_handler.py +526 -0
- revenium_middleware/openai/middleware.py +1451 -0
- revenium_middleware/openai/prompt_extractor.py +173 -0
- revenium_middleware/openai/provider.py +170 -0
- revenium_middleware/openai/summary_printer.py +292 -0
- revenium_middleware/openai/trace_fields.py +98 -0
- revenium_middleware/perplexity/__init__.py +97 -0
- revenium_middleware/perplexity/middleware.py +379 -0
- revenium_middleware/perplexity/perplexity_sdk.py +256 -0
- revenium_middleware/perplexity/provider.py +84 -0
- revenium_middleware/perplexity/trace_fields.py +25 -0
- revenium_python_sdk-0.1.0.dist-info/METADATA +252 -0
- revenium_python_sdk-0.1.0.dist-info/RECORD +73 -0
- revenium_python_sdk-0.1.0.dist-info/WHEEL +5 -0
- revenium_python_sdk-0.1.0.dist-info/licenses/LICENSE +21 -0
- revenium_python_sdk-0.1.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,1162 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Vertex AI SDK middleware for Revenium.
|
|
3
|
+
|
|
4
|
+
This module provides middleware for the native Vertex AI SDK (vertexai package),
|
|
5
|
+
offering enhanced features like comprehensive token counting and local tokenization.
|
|
6
|
+
|
|
7
|
+
Key advantages over Google AI SDK:
|
|
8
|
+
- Full token counting support including embeddings
|
|
9
|
+
- Local tokenization capabilities
|
|
10
|
+
- Enhanced metadata and usage tracking
|
|
11
|
+
- Better integration with Google Cloud services
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
import datetime
|
|
15
|
+
import logging
|
|
16
|
+
from typing import Dict, Any, Optional, List, Tuple
|
|
17
|
+
|
|
18
|
+
import wrapt
|
|
19
|
+
from revenium_middleware import run_async_in_thread
|
|
20
|
+
|
|
21
|
+
# Import common utilities and types
|
|
22
|
+
from ..common import (
|
|
23
|
+
OperationType,
|
|
24
|
+
ProviderMetadata,
|
|
25
|
+
UsageData,
|
|
26
|
+
TokenCounts,
|
|
27
|
+
normalize_stop_reason,
|
|
28
|
+
Provider,
|
|
29
|
+
create_metering_call,
|
|
30
|
+
create_image_metering_call,
|
|
31
|
+
create_video_metering_call,
|
|
32
|
+
create_usage_data,
|
|
33
|
+
extract_model_name,
|
|
34
|
+
extract_token_counts,
|
|
35
|
+
StreamingError,
|
|
36
|
+
handle_metering_error,
|
|
37
|
+
safe_getattr,
|
|
38
|
+
)
|
|
39
|
+
from ..common.trace_fields import detect_vision_content
|
|
40
|
+
|
|
41
|
+
# Vertex AI specific imports
|
|
42
|
+
from .provider import detect_provider, get_provider_metadata
|
|
43
|
+
from ..prompt_extractor import extract_prompt_data_if_enabled
|
|
44
|
+
|
|
45
|
+
logger = logging.getLogger("revenium_middleware.extension")
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def extract_vertex_ai_usage_data(
|
|
49
|
+
response: Any,
|
|
50
|
+
operation_type: OperationType,
|
|
51
|
+
request_time: datetime.datetime,
|
|
52
|
+
response_time: datetime.datetime,
|
|
53
|
+
model_name_fallback: Optional[str] = None,
|
|
54
|
+
) -> UsageData:
|
|
55
|
+
"""
|
|
56
|
+
Extract usage data from Vertex AI API responses.
|
|
57
|
+
|
|
58
|
+
This function handles the enhanced features of the Vertex AI SDK,
|
|
59
|
+
particularly the comprehensive token counting for all operations.
|
|
60
|
+
"""
|
|
61
|
+
# Get provider metadata for Vertex AI
|
|
62
|
+
provider_metadata = ProviderMetadata.for_vertex_ai_sdk()
|
|
63
|
+
|
|
64
|
+
# Extract model name - Vertex AI specific logic
|
|
65
|
+
model_name = None
|
|
66
|
+
|
|
67
|
+
# First try Vertex AI specific fields
|
|
68
|
+
if hasattr(response, "_raw_response") and response._raw_response:
|
|
69
|
+
raw_response = response._raw_response
|
|
70
|
+
if hasattr(raw_response, "model_version") and raw_response.model_version:
|
|
71
|
+
model_name = raw_response.model_version
|
|
72
|
+
logger.debug(
|
|
73
|
+
f"Extracted model name from Vertex AI _raw_response.model_version: {model_name}"
|
|
74
|
+
)
|
|
75
|
+
|
|
76
|
+
# Fallback to common extraction if not found
|
|
77
|
+
if not model_name:
|
|
78
|
+
model_name = extract_model_name(response, model_name_fallback)
|
|
79
|
+
|
|
80
|
+
# Use fallback if still not found
|
|
81
|
+
if not model_name:
|
|
82
|
+
model_name = model_name_fallback or "unknown-model"
|
|
83
|
+
|
|
84
|
+
# Clean up model name - remove Google's path prefixes
|
|
85
|
+
if model_name and isinstance(model_name, str):
|
|
86
|
+
# Remove common Google path prefixes
|
|
87
|
+
prefixes_to_remove = [
|
|
88
|
+
"publishers/google/models/",
|
|
89
|
+
"models/",
|
|
90
|
+
"google/models/",
|
|
91
|
+
"projects/",
|
|
92
|
+
]
|
|
93
|
+
for prefix in prefixes_to_remove:
|
|
94
|
+
if model_name.startswith(prefix):
|
|
95
|
+
model_name = model_name[len(prefix) :]
|
|
96
|
+
logger.debug(
|
|
97
|
+
f"Cleaned model name, removed prefix '{prefix}': {model_name}"
|
|
98
|
+
)
|
|
99
|
+
break
|
|
100
|
+
|
|
101
|
+
# Extract token counts with Vertex AI specific handling
|
|
102
|
+
if operation_type == OperationType.EMBED:
|
|
103
|
+
# Vertex AI SDK provides token counts for embeddings!
|
|
104
|
+
token_counts = extract_vertex_ai_embedding_tokens(response)
|
|
105
|
+
stop_reason = "END" # Embeddings always complete successfully
|
|
106
|
+
logger.debug(
|
|
107
|
+
f"Vertex AI embeddings token usage: {token_counts.total_tokens} tokens"
|
|
108
|
+
)
|
|
109
|
+
else: # CHAT
|
|
110
|
+
# Extract usage metadata from Vertex AI response
|
|
111
|
+
token_counts = extract_vertex_ai_generation_tokens(response)
|
|
112
|
+
|
|
113
|
+
# Determine finish reason from candidates
|
|
114
|
+
vertex_finish_reason = None
|
|
115
|
+
if hasattr(response, "candidates") and response.candidates:
|
|
116
|
+
candidate = response.candidates[0]
|
|
117
|
+
if hasattr(candidate, "finish_reason"):
|
|
118
|
+
vertex_finish_reason = candidate.finish_reason
|
|
119
|
+
logger.debug(
|
|
120
|
+
f" Raw vertex_finish_reason: {vertex_finish_reason} (type: {type(vertex_finish_reason)})"
|
|
121
|
+
)
|
|
122
|
+
|
|
123
|
+
# Convert enum to string if needed
|
|
124
|
+
if hasattr(vertex_finish_reason, "name"):
|
|
125
|
+
vertex_finish_reason = vertex_finish_reason.name
|
|
126
|
+
logger.debug(f" Converted enum to string: {vertex_finish_reason}")
|
|
127
|
+
elif not isinstance(vertex_finish_reason, str):
|
|
128
|
+
vertex_finish_reason = str(vertex_finish_reason)
|
|
129
|
+
logger.debug(f" Converted to string: {vertex_finish_reason}")
|
|
130
|
+
|
|
131
|
+
stop_reason = normalize_stop_reason(
|
|
132
|
+
vertex_finish_reason, Provider.VERTEX_AI_SDK
|
|
133
|
+
)
|
|
134
|
+
logger.debug(f" Final stop_reason after normalization: {stop_reason}")
|
|
135
|
+
logger.debug(
|
|
136
|
+
f"Vertex AI chat token usage: prompt={token_counts.input_tokens}, "
|
|
137
|
+
f"candidates={token_counts.output_tokens}, total={token_counts.total_tokens}"
|
|
138
|
+
)
|
|
139
|
+
|
|
140
|
+
# Create standardized UsageData
|
|
141
|
+
return UsageData.create(
|
|
142
|
+
operation_type=operation_type,
|
|
143
|
+
input_tokens=token_counts.input_tokens,
|
|
144
|
+
output_tokens=token_counts.output_tokens,
|
|
145
|
+
total_tokens=token_counts.total_tokens,
|
|
146
|
+
model=model_name,
|
|
147
|
+
provider_metadata=provider_metadata,
|
|
148
|
+
stop_reason=stop_reason,
|
|
149
|
+
request_time=request_time,
|
|
150
|
+
response_time=response_time,
|
|
151
|
+
cache_creation_token_count=token_counts.cached_tokens,
|
|
152
|
+
)
|
|
153
|
+
|
|
154
|
+
|
|
155
|
+
def extract_vertex_ai_generation_tokens(response: Any) -> TokenCounts:
|
|
156
|
+
"""
|
|
157
|
+
Extract token counts from Vertex AI generation responses.
|
|
158
|
+
|
|
159
|
+
Vertex AI provides comprehensive token counting in the usage_metadata.
|
|
160
|
+
"""
|
|
161
|
+
token_counts = TokenCounts(
|
|
162
|
+
input_tokens=0, output_tokens=0, total_tokens=0, cached_tokens=0
|
|
163
|
+
)
|
|
164
|
+
|
|
165
|
+
if hasattr(response, "usage_metadata") and response.usage_metadata:
|
|
166
|
+
usage_metadata = response.usage_metadata
|
|
167
|
+
|
|
168
|
+
# Vertex AI uses different attribute names than Google AI SDK
|
|
169
|
+
token_counts.input_tokens = getattr(usage_metadata, "prompt_token_count", 0)
|
|
170
|
+
token_counts.output_tokens = getattr(
|
|
171
|
+
usage_metadata, "candidates_token_count", 0
|
|
172
|
+
)
|
|
173
|
+
token_counts.total_tokens = getattr(
|
|
174
|
+
usage_metadata,
|
|
175
|
+
"total_token_count",
|
|
176
|
+
token_counts.input_tokens + token_counts.output_tokens,
|
|
177
|
+
)
|
|
178
|
+
|
|
179
|
+
# Vertex AI may provide cached token counts
|
|
180
|
+
token_counts.cached_tokens = getattr(
|
|
181
|
+
usage_metadata, "cached_content_token_count", 0
|
|
182
|
+
)
|
|
183
|
+
else:
|
|
184
|
+
logger.warning("No usage metadata found in Vertex AI generation response")
|
|
185
|
+
|
|
186
|
+
return token_counts
|
|
187
|
+
|
|
188
|
+
|
|
189
|
+
def extract_vertex_ai_embedding_tokens(response: Any) -> TokenCounts:
|
|
190
|
+
"""
|
|
191
|
+
Extract token counts from Vertex AI embedding responses.
|
|
192
|
+
|
|
193
|
+
This is a key advantage of Vertex AI SDK - embeddings include token counts!
|
|
194
|
+
"""
|
|
195
|
+
token_counts = TokenCounts(
|
|
196
|
+
input_tokens=0, output_tokens=0, total_tokens=0, cached_tokens=0
|
|
197
|
+
)
|
|
198
|
+
|
|
199
|
+
# Vertex AI embeddings response is a list of TextEmbedding objects
|
|
200
|
+
if isinstance(response, list) and len(response) > 0:
|
|
201
|
+
# Get the first embedding object
|
|
202
|
+
first_embedding = response[0]
|
|
203
|
+
|
|
204
|
+
# Check if it has statistics with token_count
|
|
205
|
+
if hasattr(first_embedding, "statistics") and first_embedding.statistics:
|
|
206
|
+
stats = first_embedding.statistics
|
|
207
|
+
if hasattr(stats, "token_count"):
|
|
208
|
+
# Convert to int if it's a float
|
|
209
|
+
token_count = (
|
|
210
|
+
int(stats.token_count)
|
|
211
|
+
if hasattr(stats.token_count, "__int__")
|
|
212
|
+
else stats.token_count
|
|
213
|
+
)
|
|
214
|
+
token_counts.input_tokens = token_count
|
|
215
|
+
token_counts.total_tokens = token_count
|
|
216
|
+
# Embeddings don't generate output tokens
|
|
217
|
+
token_counts.output_tokens = 0
|
|
218
|
+
logger.debug(
|
|
219
|
+
f"Extracted token count from Vertex AI embedding statistics: {token_count}"
|
|
220
|
+
)
|
|
221
|
+
return token_counts
|
|
222
|
+
|
|
223
|
+
# Check if the embedding has _prediction_response with metadata
|
|
224
|
+
if (
|
|
225
|
+
hasattr(first_embedding, "_prediction_response")
|
|
226
|
+
and first_embedding._prediction_response
|
|
227
|
+
):
|
|
228
|
+
pred_response = first_embedding._prediction_response
|
|
229
|
+
if hasattr(pred_response, "metadata") and pred_response.metadata:
|
|
230
|
+
# Check for billableCharacterCount or other token-related fields
|
|
231
|
+
metadata = pred_response.metadata
|
|
232
|
+
if hasattr(metadata, "billableCharacterCount"):
|
|
233
|
+
# Use billable character count as a proxy for tokens
|
|
234
|
+
char_count = metadata.billableCharacterCount
|
|
235
|
+
# Rough approximation: 4 characters per token (common for many tokenizers)
|
|
236
|
+
estimated_tokens = max(1, int(char_count / 4))
|
|
237
|
+
token_counts.input_tokens = estimated_tokens
|
|
238
|
+
token_counts.total_tokens = estimated_tokens
|
|
239
|
+
token_counts.output_tokens = 0
|
|
240
|
+
logger.debug(
|
|
241
|
+
f"Estimated token count from billable characters: {char_count} chars -> {estimated_tokens} tokens"
|
|
242
|
+
)
|
|
243
|
+
return token_counts
|
|
244
|
+
|
|
245
|
+
# Fallback: check if response itself has statistics or usage_metadata
|
|
246
|
+
elif hasattr(response, "statistics") and response.statistics:
|
|
247
|
+
# Some Vertex AI embedding responses have statistics
|
|
248
|
+
stats = response.statistics
|
|
249
|
+
if hasattr(stats, "token_count"):
|
|
250
|
+
token_counts.input_tokens = stats.token_count
|
|
251
|
+
token_counts.total_tokens = stats.token_count
|
|
252
|
+
# Embeddings don't generate output tokens
|
|
253
|
+
token_counts.output_tokens = 0
|
|
254
|
+
elif hasattr(response, "usage_metadata") and response.usage_metadata:
|
|
255
|
+
# Alternative location for token counts
|
|
256
|
+
usage_metadata = response.usage_metadata
|
|
257
|
+
token_counts.input_tokens = getattr(usage_metadata, "prompt_token_count", 0)
|
|
258
|
+
token_counts.total_tokens = getattr(
|
|
259
|
+
usage_metadata, "total_token_count", token_counts.input_tokens
|
|
260
|
+
)
|
|
261
|
+
token_counts.output_tokens = 0 # Embeddings don't generate output
|
|
262
|
+
else:
|
|
263
|
+
# If no token counts available, log warning
|
|
264
|
+
logger.debug("No token counts found in Vertex AI embedding response")
|
|
265
|
+
|
|
266
|
+
return token_counts
|
|
267
|
+
|
|
268
|
+
|
|
269
|
+
def create_vertex_ai_metering_call(
|
|
270
|
+
response: Any,
|
|
271
|
+
operation_type: OperationType,
|
|
272
|
+
request_time_dt: datetime.datetime,
|
|
273
|
+
usage_metadata: Dict[str, Any],
|
|
274
|
+
time_to_first_token: int = 0,
|
|
275
|
+
is_streamed: bool = False,
|
|
276
|
+
model_name_fallback: Optional[str] = None,
|
|
277
|
+
# Prompt capture fields
|
|
278
|
+
system_prompt: Optional[str] = None,
|
|
279
|
+
input_messages: Optional[str] = None,
|
|
280
|
+
output_response: Optional[str] = None,
|
|
281
|
+
prompts_truncated: Optional[bool] = None,
|
|
282
|
+
) -> None:
|
|
283
|
+
"""
|
|
284
|
+
Create and execute a metering call for Vertex AI SDK responses.
|
|
285
|
+
|
|
286
|
+
This is the main function used by the wrapper functions.
|
|
287
|
+
"""
|
|
288
|
+
# Record response timing
|
|
289
|
+
response_time_dt = datetime.datetime.now(datetime.timezone.utc)
|
|
290
|
+
|
|
291
|
+
# Extract usage data using Vertex AI specific logic
|
|
292
|
+
usage_data = extract_vertex_ai_usage_data(
|
|
293
|
+
response=response,
|
|
294
|
+
operation_type=operation_type,
|
|
295
|
+
request_time=request_time_dt,
|
|
296
|
+
response_time=response_time_dt,
|
|
297
|
+
model_name_fallback=model_name_fallback,
|
|
298
|
+
)
|
|
299
|
+
|
|
300
|
+
# Create metering call using common utilities
|
|
301
|
+
create_metering_call(
|
|
302
|
+
usage_data=usage_data,
|
|
303
|
+
usage_metadata=usage_metadata,
|
|
304
|
+
time_to_first_token=time_to_first_token,
|
|
305
|
+
is_streamed=is_streamed,
|
|
306
|
+
# Prompt capture fields
|
|
307
|
+
system_prompt=system_prompt,
|
|
308
|
+
input_messages=input_messages,
|
|
309
|
+
output_response=output_response,
|
|
310
|
+
prompts_truncated=prompts_truncated,
|
|
311
|
+
)
|
|
312
|
+
|
|
313
|
+
|
|
314
|
+
# Dynamic wrapper discovery and application for Vertex AI GenerativeModel.generate_content
|
|
315
|
+
def _apply_generate_content_wrappers():
|
|
316
|
+
"""
|
|
317
|
+
Dynamically discover and wrap all Vertex AI GenerativeModel.generate_content methods.
|
|
318
|
+
This handles current and future module path variations like:
|
|
319
|
+
- vertexai.generative_models.GenerativeModel
|
|
320
|
+
- vertexai.preview.generative_models.GenerativeModel
|
|
321
|
+
- vertexai.v1.generative_models.GenerativeModel
|
|
322
|
+
- etc.
|
|
323
|
+
"""
|
|
324
|
+
import sys
|
|
325
|
+
import importlib
|
|
326
|
+
|
|
327
|
+
# Known module patterns to try
|
|
328
|
+
module_patterns = [
|
|
329
|
+
"vertexai.generative_models",
|
|
330
|
+
"vertexai.preview.generative_models",
|
|
331
|
+
"vertexai.v1.generative_models",
|
|
332
|
+
"vertexai.v1beta1.generative_models",
|
|
333
|
+
"vertexai.v2.generative_models",
|
|
334
|
+
"vertexai.beta.generative_models",
|
|
335
|
+
"vertexai.alpha.generative_models",
|
|
336
|
+
]
|
|
337
|
+
|
|
338
|
+
wrapped_modules = []
|
|
339
|
+
|
|
340
|
+
for module_path in module_patterns:
|
|
341
|
+
try:
|
|
342
|
+
# Try to import the module
|
|
343
|
+
module = importlib.import_module(module_path)
|
|
344
|
+
|
|
345
|
+
# Check if GenerativeModel class exists
|
|
346
|
+
if hasattr(module, "GenerativeModel"):
|
|
347
|
+
generative_model_class = getattr(module, "GenerativeModel")
|
|
348
|
+
|
|
349
|
+
# Check if generate_content method exists
|
|
350
|
+
if hasattr(generative_model_class, "generate_content"):
|
|
351
|
+
logger.debug(
|
|
352
|
+
f"Found GenerativeModel.generate_content in {module_path}"
|
|
353
|
+
)
|
|
354
|
+
|
|
355
|
+
# Apply wrapper using wrapt
|
|
356
|
+
@wrapt.patch_function_wrapper(
|
|
357
|
+
module_path, "GenerativeModel.generate_content"
|
|
358
|
+
)
|
|
359
|
+
def generate_content_wrapper_dynamic(
|
|
360
|
+
wrapped, instance, args, kwargs
|
|
361
|
+
):
|
|
362
|
+
return generate_content_wrapper_impl(
|
|
363
|
+
wrapped, instance, args, kwargs
|
|
364
|
+
)
|
|
365
|
+
|
|
366
|
+
wrapped_modules.append(module_path)
|
|
367
|
+
logger.debug(
|
|
368
|
+
f" Applied wrapper to {module_path}.GenerativeModel.generate_content"
|
|
369
|
+
)
|
|
370
|
+
else:
|
|
371
|
+
logger.debug(
|
|
372
|
+
f" {module_path}.GenerativeModel exists but no generate_content method"
|
|
373
|
+
)
|
|
374
|
+
else:
|
|
375
|
+
logger.debug(f" {module_path} exists but no GenerativeModel class")
|
|
376
|
+
|
|
377
|
+
except ImportError:
|
|
378
|
+
logger.debug(f" Module {module_path} not available")
|
|
379
|
+
except Exception as e:
|
|
380
|
+
logger.debug(f" Error checking {module_path}: {e}")
|
|
381
|
+
|
|
382
|
+
if wrapped_modules:
|
|
383
|
+
logger.info(
|
|
384
|
+
f" Vertex AI GenerativeModel wrappers applied to: {', '.join(wrapped_modules)}"
|
|
385
|
+
)
|
|
386
|
+
else:
|
|
387
|
+
logger.warning(" No Vertex AI GenerativeModel modules found to wrap")
|
|
388
|
+
|
|
389
|
+
return wrapped_modules
|
|
390
|
+
|
|
391
|
+
|
|
392
|
+
def generate_content_wrapper_impl(wrapped, instance, args, kwargs):
|
|
393
|
+
"""Enhanced wrapper that handles both streaming and non-streaming Vertex AI calls."""
|
|
394
|
+
logger.debug("Enhanced Vertex AI generate_content wrapper called!")
|
|
395
|
+
logger.debug(f"Wrapper args: {args}")
|
|
396
|
+
logger.debug(f"Wrapper kwargs: {kwargs}")
|
|
397
|
+
logger.debug(f"Instance type: {type(instance)}")
|
|
398
|
+
|
|
399
|
+
# Extract usage metadata from instance or kwargs
|
|
400
|
+
usage_metadata = getattr(instance, "_revenium_usage_metadata", {}) or kwargs.pop(
|
|
401
|
+
"usage_metadata", {}
|
|
402
|
+
)
|
|
403
|
+
logger.debug(f"Captured usage metadata for generate_content: {usage_metadata}")
|
|
404
|
+
logger.debug(
|
|
405
|
+
f"Instance has _revenium_usage_metadata: {hasattr(instance, '_revenium_usage_metadata')}"
|
|
406
|
+
)
|
|
407
|
+
if hasattr(instance, "_revenium_usage_metadata"):
|
|
408
|
+
logger.debug(
|
|
409
|
+
f"Instance._revenium_usage_metadata value: {getattr(instance, '_revenium_usage_metadata')}"
|
|
410
|
+
)
|
|
411
|
+
|
|
412
|
+
# Try to extract model name from the instance
|
|
413
|
+
model_name_from_instance = None
|
|
414
|
+
for attr in [
|
|
415
|
+
"_model_name",
|
|
416
|
+
"model_name",
|
|
417
|
+
"_model_id",
|
|
418
|
+
"model_id",
|
|
419
|
+
"_model",
|
|
420
|
+
"model",
|
|
421
|
+
]:
|
|
422
|
+
if hasattr(instance, attr):
|
|
423
|
+
model_name_from_instance = getattr(instance, attr)
|
|
424
|
+
logger.debug(
|
|
425
|
+
f"Found model name in instance.{attr}: {model_name_from_instance}"
|
|
426
|
+
)
|
|
427
|
+
break
|
|
428
|
+
|
|
429
|
+
# Clean up the instance model name too
|
|
430
|
+
if model_name_from_instance and isinstance(model_name_from_instance, str):
|
|
431
|
+
# Remove common Google path prefixes
|
|
432
|
+
prefixes_to_remove = [
|
|
433
|
+
"publishers/google/models/",
|
|
434
|
+
"models/",
|
|
435
|
+
"google/models/",
|
|
436
|
+
"projects/",
|
|
437
|
+
]
|
|
438
|
+
for prefix in prefixes_to_remove:
|
|
439
|
+
if model_name_from_instance.startswith(prefix):
|
|
440
|
+
model_name_from_instance = model_name_from_instance[len(prefix) :]
|
|
441
|
+
logger.debug(
|
|
442
|
+
f"Cleaned instance model name, removed prefix '{prefix}': {model_name_from_instance}"
|
|
443
|
+
)
|
|
444
|
+
break
|
|
445
|
+
|
|
446
|
+
if not model_name_from_instance:
|
|
447
|
+
logger.debug(
|
|
448
|
+
f"Could not find model name in instance. Available attributes: {dir(instance)}"
|
|
449
|
+
)
|
|
450
|
+
# Try to get it from the instance string representation
|
|
451
|
+
instance_str = str(instance)
|
|
452
|
+
if "model_name=" in instance_str:
|
|
453
|
+
# Extract from string like "GenerativeModel(model_name='gemini-2.0-flash-lite-001')"
|
|
454
|
+
import re
|
|
455
|
+
|
|
456
|
+
match = re.search(r"model_name='([^']+)'", instance_str)
|
|
457
|
+
if match:
|
|
458
|
+
model_name_from_instance = match.group(1)
|
|
459
|
+
logger.debug(
|
|
460
|
+
f"Extracted model name from instance string: {model_name_from_instance}"
|
|
461
|
+
)
|
|
462
|
+
elif "models/" in instance_str:
|
|
463
|
+
# Extract from string like "models/gemini-2.0-flash-lite-001"
|
|
464
|
+
import re
|
|
465
|
+
|
|
466
|
+
match = re.search(r"models/([^'\s)]+)", instance_str)
|
|
467
|
+
if match:
|
|
468
|
+
model_name_from_instance = match.group(1)
|
|
469
|
+
logger.debug(
|
|
470
|
+
f"Extracted model name from instance string (models/): {model_name_from_instance}"
|
|
471
|
+
)
|
|
472
|
+
|
|
473
|
+
# Detect vision content in the request
|
|
474
|
+
# Vertex AI generate_content takes contents as first positional arg or 'contents' kwarg
|
|
475
|
+
contents = kwargs.get("contents") or (args[0] if args else None)
|
|
476
|
+
has_vision = detect_vision_content(contents)
|
|
477
|
+
if has_vision:
|
|
478
|
+
usage_metadata["has_vision_content"] = True
|
|
479
|
+
logger.debug("Vision content detected in Vertex AI generate_content request")
|
|
480
|
+
|
|
481
|
+
# Check if this is a streaming call
|
|
482
|
+
is_streaming = kwargs.get("stream", False)
|
|
483
|
+
|
|
484
|
+
# Store kwargs and args for prompt extraction
|
|
485
|
+
request_kwargs = kwargs.copy()
|
|
486
|
+
request_args = args
|
|
487
|
+
|
|
488
|
+
# Record request time
|
|
489
|
+
request_time_dt = datetime.datetime.now(datetime.timezone.utc)
|
|
490
|
+
logger.debug(
|
|
491
|
+
f"Calling wrapped Vertex AI generate_content function (streaming={is_streaming}) with args: {args}, kwargs: {kwargs}"
|
|
492
|
+
)
|
|
493
|
+
|
|
494
|
+
# Call the original Vertex AI function
|
|
495
|
+
response = wrapped(*args, **kwargs)
|
|
496
|
+
|
|
497
|
+
if is_streaming:
|
|
498
|
+
logger.debug("Handling Vertex AI streaming response")
|
|
499
|
+
# Return wrapped stream that will meter usage when complete
|
|
500
|
+
return handle_vertex_ai_streaming_response(
|
|
501
|
+
stream=response,
|
|
502
|
+
request_time_dt=request_time_dt,
|
|
503
|
+
usage_metadata=usage_metadata,
|
|
504
|
+
model_name_fallback=model_name_from_instance,
|
|
505
|
+
request_kwargs=request_kwargs,
|
|
506
|
+
request_args=request_args,
|
|
507
|
+
)
|
|
508
|
+
else:
|
|
509
|
+
logger.debug("Handling Vertex AI non-streaming response: %s", response)
|
|
510
|
+
|
|
511
|
+
# Extract prompt data if capture is enabled
|
|
512
|
+
system_prompt, input_messages, output_response, prompts_truncated = (
|
|
513
|
+
extract_prompt_data_if_enabled(request_kwargs, args=request_args, response=response)
|
|
514
|
+
)
|
|
515
|
+
|
|
516
|
+
# Handle non-streaming response immediately
|
|
517
|
+
create_vertex_ai_metering_call(
|
|
518
|
+
response=response,
|
|
519
|
+
operation_type=OperationType.CHAT,
|
|
520
|
+
request_time_dt=request_time_dt,
|
|
521
|
+
usage_metadata=usage_metadata,
|
|
522
|
+
model_name_fallback=model_name_from_instance,
|
|
523
|
+
# Prompt capture fields
|
|
524
|
+
system_prompt=system_prompt,
|
|
525
|
+
input_messages=input_messages,
|
|
526
|
+
output_response=output_response,
|
|
527
|
+
prompts_truncated=prompts_truncated,
|
|
528
|
+
)
|
|
529
|
+
return response
|
|
530
|
+
|
|
531
|
+
|
|
532
|
+
# Wrapper for Vertex AI TextEmbeddingModel.get_embeddings method
|
|
533
|
+
@wrapt.patch_function_wrapper(
|
|
534
|
+
"vertexai.language_models", "TextEmbeddingModel.get_embeddings"
|
|
535
|
+
)
|
|
536
|
+
def get_embeddings_wrapper(wrapped, instance, args, kwargs):
|
|
537
|
+
"""Wraps the vertexai.language_models.TextEmbeddingModel.get_embeddings method to log token usage."""
|
|
538
|
+
logger.debug("Vertex AI get_embeddings wrapper called")
|
|
539
|
+
|
|
540
|
+
# Extract usage metadata from instance or kwargs
|
|
541
|
+
usage_metadata = getattr(instance, "_revenium_usage_metadata", {}) or kwargs.pop(
|
|
542
|
+
"usage_metadata", {}
|
|
543
|
+
)
|
|
544
|
+
|
|
545
|
+
# Try to extract model name from the instance using the same logic as generate_content
|
|
546
|
+
model_name_from_instance = None
|
|
547
|
+
for attr in [
|
|
548
|
+
"_model_name",
|
|
549
|
+
"model_name",
|
|
550
|
+
"_model_id",
|
|
551
|
+
"model_id",
|
|
552
|
+
"_model",
|
|
553
|
+
"model",
|
|
554
|
+
]:
|
|
555
|
+
if hasattr(instance, attr):
|
|
556
|
+
model_name_from_instance = getattr(instance, attr)
|
|
557
|
+
logger.debug(
|
|
558
|
+
f"Found model name in embeddings instance.{attr}: {model_name_from_instance}"
|
|
559
|
+
)
|
|
560
|
+
break
|
|
561
|
+
|
|
562
|
+
if not model_name_from_instance:
|
|
563
|
+
logger.debug(
|
|
564
|
+
f"Could not find model name in embeddings instance. Available attributes: {dir(instance)}"
|
|
565
|
+
)
|
|
566
|
+
# Try to get it from the instance string representation
|
|
567
|
+
instance_str = str(instance)
|
|
568
|
+
if "model_name=" in instance_str:
|
|
569
|
+
# Extract from string like "TextEmbeddingModel(model_name='text-embedding-004')"
|
|
570
|
+
import re
|
|
571
|
+
|
|
572
|
+
match = re.search(r"model_name='([^']+)'", instance_str)
|
|
573
|
+
if match:
|
|
574
|
+
model_name_from_instance = match.group(1)
|
|
575
|
+
logger.debug(
|
|
576
|
+
f"Extracted model name from embeddings instance string: {model_name_from_instance}"
|
|
577
|
+
)
|
|
578
|
+
elif "models/" in instance_str:
|
|
579
|
+
# Extract from string like "models/text-embedding-004"
|
|
580
|
+
import re
|
|
581
|
+
|
|
582
|
+
match = re.search(r"models/([^'\s)]+)", instance_str)
|
|
583
|
+
if match:
|
|
584
|
+
model_name_from_instance = match.group(1)
|
|
585
|
+
logger.debug(
|
|
586
|
+
f"Extracted model name from embeddings instance string (models/): {model_name_from_instance}"
|
|
587
|
+
)
|
|
588
|
+
|
|
589
|
+
logger.debug(
|
|
590
|
+
f"Final captured model name from Vertex AI embeddings instance: {model_name_from_instance}"
|
|
591
|
+
)
|
|
592
|
+
|
|
593
|
+
# Record request time
|
|
594
|
+
request_time_dt = datetime.datetime.now(datetime.timezone.utc)
|
|
595
|
+
logger.debug(
|
|
596
|
+
f"Calling wrapped Vertex AI get_embeddings function with args: {args}, kwargs: {kwargs}"
|
|
597
|
+
)
|
|
598
|
+
|
|
599
|
+
# Call the original Vertex AI function
|
|
600
|
+
response = wrapped(*args, **kwargs)
|
|
601
|
+
|
|
602
|
+
logger.debug("Handling Vertex AI get_embeddings response: %s", response)
|
|
603
|
+
|
|
604
|
+
# Create metering call for embeddings
|
|
605
|
+
create_vertex_ai_metering_call(
|
|
606
|
+
response=response,
|
|
607
|
+
operation_type=OperationType.EMBED,
|
|
608
|
+
request_time_dt=request_time_dt,
|
|
609
|
+
usage_metadata=usage_metadata,
|
|
610
|
+
model_name_fallback=model_name_from_instance,
|
|
611
|
+
)
|
|
612
|
+
|
|
613
|
+
return response
|
|
614
|
+
|
|
615
|
+
|
|
616
|
+
def handle_vertex_ai_streaming_response(
|
|
617
|
+
stream, request_time_dt, usage_metadata, model_name_fallback=None, request_kwargs=None, request_args=None
|
|
618
|
+
):
|
|
619
|
+
"""
|
|
620
|
+
Handle streaming responses from Vertex AI.
|
|
621
|
+
Wraps the stream to collect metrics and log them after completion.
|
|
622
|
+
"""
|
|
623
|
+
|
|
624
|
+
class VertexAIStreamWrapper:
|
|
625
|
+
def __init__(self, stream):
|
|
626
|
+
self.stream = stream
|
|
627
|
+
self.chunks = []
|
|
628
|
+
self.accumulated_text = [] # For prompt capture
|
|
629
|
+
self.model = model_name_fallback
|
|
630
|
+
self.finish_reason = None
|
|
631
|
+
self.usage_metadata = None
|
|
632
|
+
self.first_chunk_time = None
|
|
633
|
+
self._closed = False
|
|
634
|
+
self._usage_logged = False
|
|
635
|
+
self.streaming_truncated = False # Track if streaming response was truncated
|
|
636
|
+
|
|
637
|
+
# Limit chunk storage to prevent memory issues
|
|
638
|
+
self._max_chunks = 1000
|
|
639
|
+
|
|
640
|
+
def __iter__(self):
|
|
641
|
+
return self
|
|
642
|
+
|
|
643
|
+
def __next__(self):
|
|
644
|
+
if self._closed:
|
|
645
|
+
raise StopIteration("Stream has been closed")
|
|
646
|
+
|
|
647
|
+
try:
|
|
648
|
+
chunk = next(self.stream)
|
|
649
|
+
self._process_chunk(chunk)
|
|
650
|
+
return chunk
|
|
651
|
+
except StopIteration:
|
|
652
|
+
self._finalize()
|
|
653
|
+
raise
|
|
654
|
+
except Exception as e:
|
|
655
|
+
self._handle_error(e)
|
|
656
|
+
raise
|
|
657
|
+
|
|
658
|
+
def __enter__(self):
|
|
659
|
+
return self
|
|
660
|
+
|
|
661
|
+
def __exit__(self, exc_type, exc_val, exc_tb):
|
|
662
|
+
self.close()
|
|
663
|
+
return False # Don't suppress exceptions
|
|
664
|
+
|
|
665
|
+
def close(self):
|
|
666
|
+
"""Properly close the stream and clean up resources."""
|
|
667
|
+
if not self._closed:
|
|
668
|
+
self._closed = True
|
|
669
|
+
if not self._usage_logged:
|
|
670
|
+
try:
|
|
671
|
+
self._log_usage()
|
|
672
|
+
except Exception as e:
|
|
673
|
+
logger.error(
|
|
674
|
+
"Error logging usage during Vertex AI stream cleanup: %s", e
|
|
675
|
+
)
|
|
676
|
+
|
|
677
|
+
# Clear chunks to free memory
|
|
678
|
+
self.chunks.clear()
|
|
679
|
+
|
|
680
|
+
# Close underlying stream if it has a close method
|
|
681
|
+
if hasattr(self.stream, "close"):
|
|
682
|
+
try:
|
|
683
|
+
self.stream.close()
|
|
684
|
+
except Exception as e:
|
|
685
|
+
logger.debug("Error closing underlying Vertex AI stream: %s", e)
|
|
686
|
+
|
|
687
|
+
def _finalize(self):
|
|
688
|
+
"""Finalize the stream and log usage."""
|
|
689
|
+
if not self._usage_logged:
|
|
690
|
+
self._log_usage()
|
|
691
|
+
self._usage_logged = True
|
|
692
|
+
|
|
693
|
+
def _handle_error(self, error: Exception):
|
|
694
|
+
"""Handle errors during streaming."""
|
|
695
|
+
logger.error("Error in Vertex AI streaming response: %s", error)
|
|
696
|
+
if not self._usage_logged:
|
|
697
|
+
# Try to log partial usage data
|
|
698
|
+
try:
|
|
699
|
+
self._log_usage()
|
|
700
|
+
self._usage_logged = True
|
|
701
|
+
except Exception as log_error:
|
|
702
|
+
logger.error(
|
|
703
|
+
"Failed to log Vertex AI usage after stream error: %s",
|
|
704
|
+
log_error,
|
|
705
|
+
)
|
|
706
|
+
|
|
707
|
+
def _process_chunk(self, chunk):
|
|
708
|
+
"""Process each chunk to extract metadata"""
|
|
709
|
+
# Limit chunk storage to prevent memory issues
|
|
710
|
+
if len(self.chunks) < self._max_chunks:
|
|
711
|
+
self.chunks.append(chunk)
|
|
712
|
+
elif len(self.chunks) == self._max_chunks:
|
|
713
|
+
logger.warning(
|
|
714
|
+
"Reached maximum chunk limit (%d) for Vertex AI stream, not storing additional chunks",
|
|
715
|
+
self._max_chunks,
|
|
716
|
+
)
|
|
717
|
+
|
|
718
|
+
# Record time of first chunk
|
|
719
|
+
if self.first_chunk_time is None:
|
|
720
|
+
self.first_chunk_time = datetime.datetime.now(datetime.timezone.utc)
|
|
721
|
+
|
|
722
|
+
# Extract model name from chunk if available using safe access
|
|
723
|
+
if self.model is None:
|
|
724
|
+
self.model = extract_model_name(chunk, self.model)
|
|
725
|
+
|
|
726
|
+
# Accumulate text for prompt capture (with early truncation to prevent unbounded memory growth)
|
|
727
|
+
from ..config import Config
|
|
728
|
+
current_len = sum(len(t) for t in self.accumulated_text)
|
|
729
|
+
|
|
730
|
+
if hasattr(chunk, 'text') and chunk.text:
|
|
731
|
+
# Check if adding this chunk would exceed the limit
|
|
732
|
+
chunk_len = len(chunk.text)
|
|
733
|
+
if current_len + chunk_len <= Config.MAX_PROMPT_LENGTH:
|
|
734
|
+
self.accumulated_text.append(chunk.text)
|
|
735
|
+
elif current_len < Config.MAX_PROMPT_LENGTH:
|
|
736
|
+
# Partial append: only add what fits
|
|
737
|
+
remaining = Config.MAX_PROMPT_LENGTH - current_len
|
|
738
|
+
self.accumulated_text.append(chunk.text[:remaining])
|
|
739
|
+
self.streaming_truncated = True
|
|
740
|
+
else:
|
|
741
|
+
# Already at limit, mark as truncated
|
|
742
|
+
self.streaming_truncated = True
|
|
743
|
+
elif hasattr(chunk, 'candidates') and chunk.candidates:
|
|
744
|
+
for candidate in chunk.candidates:
|
|
745
|
+
if hasattr(candidate, 'content') and candidate.content:
|
|
746
|
+
if hasattr(candidate.content, 'parts'):
|
|
747
|
+
for part in candidate.content.parts:
|
|
748
|
+
if hasattr(part, 'text') and part.text:
|
|
749
|
+
part_len = len(part.text)
|
|
750
|
+
if current_len + part_len <= Config.MAX_PROMPT_LENGTH:
|
|
751
|
+
self.accumulated_text.append(part.text)
|
|
752
|
+
current_len += part_len
|
|
753
|
+
elif current_len < Config.MAX_PROMPT_LENGTH:
|
|
754
|
+
# Partial append: only add what fits
|
|
755
|
+
remaining = Config.MAX_PROMPT_LENGTH - current_len
|
|
756
|
+
self.accumulated_text.append(part.text[:remaining])
|
|
757
|
+
current_len = Config.MAX_PROMPT_LENGTH
|
|
758
|
+
self.streaming_truncated = True
|
|
759
|
+
break
|
|
760
|
+
else:
|
|
761
|
+
# Already at limit
|
|
762
|
+
self.streaming_truncated = True
|
|
763
|
+
break
|
|
764
|
+
|
|
765
|
+
# Check for finish reason and usage metadata in the chunk using safe access
|
|
766
|
+
candidates = safe_getattr(chunk, "candidates")
|
|
767
|
+
if candidates and len(candidates) > 0:
|
|
768
|
+
candidate = candidates[0]
|
|
769
|
+
finish_reason = safe_getattr(candidate, "finish_reason")
|
|
770
|
+
if finish_reason:
|
|
771
|
+
self.finish_reason = finish_reason
|
|
772
|
+
|
|
773
|
+
# Check for usage metadata in the chunk (final chunk typically has this)
|
|
774
|
+
usage_metadata = safe_getattr(chunk, "usage_metadata")
|
|
775
|
+
if usage_metadata:
|
|
776
|
+
self.usage_metadata = usage_metadata
|
|
777
|
+
|
|
778
|
+
def _log_usage(self):
|
|
779
|
+
"""Log usage after stream completion"""
|
|
780
|
+
try:
|
|
781
|
+
if not self.chunks:
|
|
782
|
+
logger.warning("No chunks received in Vertex AI streaming response")
|
|
783
|
+
return
|
|
784
|
+
|
|
785
|
+
# Calculate time to first token
|
|
786
|
+
time_to_first_token = 0
|
|
787
|
+
if self.first_chunk_time:
|
|
788
|
+
time_to_first_token = int(
|
|
789
|
+
(self.first_chunk_time - request_time_dt).total_seconds() * 1000
|
|
790
|
+
)
|
|
791
|
+
|
|
792
|
+
# Extract prompt data if capture is enabled
|
|
793
|
+
accumulated_content = ''.join(self.accumulated_text) if self.accumulated_text else None
|
|
794
|
+
# Append truncation marker if streaming was truncated
|
|
795
|
+
if self.streaming_truncated and accumulated_content:
|
|
796
|
+
accumulated_content += "...[TRUNCATED]"
|
|
797
|
+
|
|
798
|
+
system_prompt, input_messages, output_response, prompts_truncated = (
|
|
799
|
+
extract_prompt_data_if_enabled(
|
|
800
|
+
request_kwargs or {},
|
|
801
|
+
args=request_args,
|
|
802
|
+
accumulated_content=accumulated_content
|
|
803
|
+
)
|
|
804
|
+
)
|
|
805
|
+
|
|
806
|
+
# Update truncation flag if streaming was truncated
|
|
807
|
+
if self.streaming_truncated:
|
|
808
|
+
prompts_truncated = True
|
|
809
|
+
|
|
810
|
+
# Create a synthetic response object for usage extraction
|
|
811
|
+
class SyntheticResponse:
|
|
812
|
+
def __init__(self, model_name, usage_metadata, candidates):
|
|
813
|
+
self.model_name = model_name
|
|
814
|
+
self.usage_metadata = usage_metadata
|
|
815
|
+
self.candidates = candidates
|
|
816
|
+
|
|
817
|
+
# Create synthetic response from collected data
|
|
818
|
+
synthetic_response = SyntheticResponse(
|
|
819
|
+
model_name=self.model,
|
|
820
|
+
usage_metadata=self.usage_metadata,
|
|
821
|
+
candidates=(
|
|
822
|
+
[
|
|
823
|
+
type(
|
|
824
|
+
"obj", (object,), {"finish_reason": self.finish_reason}
|
|
825
|
+
)()
|
|
826
|
+
]
|
|
827
|
+
if self.finish_reason
|
|
828
|
+
else []
|
|
829
|
+
),
|
|
830
|
+
)
|
|
831
|
+
|
|
832
|
+
# Create metering call for streaming response
|
|
833
|
+
create_vertex_ai_metering_call(
|
|
834
|
+
response=synthetic_response,
|
|
835
|
+
operation_type=OperationType.CHAT,
|
|
836
|
+
request_time_dt=request_time_dt,
|
|
837
|
+
usage_metadata=usage_metadata,
|
|
838
|
+
time_to_first_token=time_to_first_token,
|
|
839
|
+
is_streamed=True,
|
|
840
|
+
model_name_fallback=self.model,
|
|
841
|
+
# Prompt capture fields
|
|
842
|
+
system_prompt=system_prompt,
|
|
843
|
+
input_messages=input_messages,
|
|
844
|
+
output_response=output_response,
|
|
845
|
+
prompts_truncated=prompts_truncated,
|
|
846
|
+
)
|
|
847
|
+
|
|
848
|
+
logger.debug(
|
|
849
|
+
"Vertex AI streaming usage logged: model=%s, chunks=%d, time_to_first_token=%dms",
|
|
850
|
+
self.model,
|
|
851
|
+
len(self.chunks),
|
|
852
|
+
time_to_first_token,
|
|
853
|
+
)
|
|
854
|
+
|
|
855
|
+
except Exception as e:
|
|
856
|
+
# Don't let logging errors break the stream
|
|
857
|
+
logger.error("Error logging Vertex AI streaming usage: %s", e)
|
|
858
|
+
raise StreamingError(
|
|
859
|
+
f"Failed to log Vertex AI streaming usage: {str(e)}",
|
|
860
|
+
chunk_count=len(self.chunks) if self.chunks else 0,
|
|
861
|
+
stream_state="completed",
|
|
862
|
+
) from e
|
|
863
|
+
|
|
864
|
+
return VertexAIStreamWrapper(stream)
|
|
865
|
+
|
|
866
|
+
|
|
867
|
+
# --- Vertex AI ImageGenerationModel wrapper (Imagen) ---
|
|
868
|
+
|
|
869
|
+
def _apply_imagen_wrappers():
|
|
870
|
+
"""
|
|
871
|
+
Dynamically discover and wrap Vertex AI ImageGenerationModel.generate_images.
|
|
872
|
+
Handles multiple module paths for forward compatibility.
|
|
873
|
+
"""
|
|
874
|
+
import importlib
|
|
875
|
+
|
|
876
|
+
module_patterns = [
|
|
877
|
+
"vertexai.preview.vision_models",
|
|
878
|
+
"vertexai.vision_models",
|
|
879
|
+
]
|
|
880
|
+
|
|
881
|
+
wrapped_modules = []
|
|
882
|
+
|
|
883
|
+
for module_path in module_patterns:
|
|
884
|
+
try:
|
|
885
|
+
module = importlib.import_module(module_path)
|
|
886
|
+
|
|
887
|
+
if hasattr(module, "ImageGenerationModel"):
|
|
888
|
+
img_model_class = getattr(module, "ImageGenerationModel")
|
|
889
|
+
|
|
890
|
+
if hasattr(img_model_class, "generate_images"):
|
|
891
|
+
@wrapt.patch_function_wrapper(
|
|
892
|
+
module_path, "ImageGenerationModel.generate_images"
|
|
893
|
+
)
|
|
894
|
+
def generate_images_wrapper_dynamic(wrapped, instance, args, kwargs):
|
|
895
|
+
return _imagen_generate_images_impl(wrapped, instance, args, kwargs)
|
|
896
|
+
|
|
897
|
+
wrapped_modules.append(module_path + ".ImageGenerationModel.generate_images")
|
|
898
|
+
logger.debug(f" Applied Imagen wrapper to {module_path}.ImageGenerationModel.generate_images")
|
|
899
|
+
|
|
900
|
+
if hasattr(img_model_class, "edit_image"):
|
|
901
|
+
@wrapt.patch_function_wrapper(
|
|
902
|
+
module_path, "ImageGenerationModel.edit_image"
|
|
903
|
+
)
|
|
904
|
+
def edit_image_wrapper_dynamic(wrapped, instance, args, kwargs):
|
|
905
|
+
return _imagen_edit_image_impl(wrapped, instance, args, kwargs)
|
|
906
|
+
|
|
907
|
+
wrapped_modules.append(module_path + ".ImageGenerationModel.edit_image")
|
|
908
|
+
logger.debug(f" Applied Imagen wrapper to {module_path}.ImageGenerationModel.edit_image")
|
|
909
|
+
|
|
910
|
+
except ImportError:
|
|
911
|
+
logger.debug(f" Module {module_path} not available for Imagen")
|
|
912
|
+
except Exception as e:
|
|
913
|
+
logger.debug(f" Error applying Imagen wrapper to {module_path}: {e}")
|
|
914
|
+
|
|
915
|
+
if wrapped_modules:
|
|
916
|
+
logger.info(f" Vertex AI Imagen wrappers applied to: {', '.join(wrapped_modules)}")
|
|
917
|
+
|
|
918
|
+
return wrapped_modules
|
|
919
|
+
|
|
920
|
+
|
|
921
|
+
def _extract_model_name_from_instance(instance) -> Optional[str]:
|
|
922
|
+
"""Extract and clean model name from a Vertex AI model instance."""
|
|
923
|
+
model_name = None
|
|
924
|
+
for attr in ["_model_id", "model_id", "_model_name", "model_name", "_model", "model"]:
|
|
925
|
+
if hasattr(instance, attr):
|
|
926
|
+
model_name = getattr(instance, attr)
|
|
927
|
+
if model_name:
|
|
928
|
+
break
|
|
929
|
+
|
|
930
|
+
if model_name and isinstance(model_name, str):
|
|
931
|
+
prefixes = ["publishers/google/models/", "models/", "google/models/", "projects/"]
|
|
932
|
+
for prefix in prefixes:
|
|
933
|
+
if model_name.startswith(prefix):
|
|
934
|
+
model_name = model_name[len(prefix):]
|
|
935
|
+
break
|
|
936
|
+
|
|
937
|
+
return model_name
|
|
938
|
+
|
|
939
|
+
|
|
940
|
+
def _imagen_generate_images_impl(wrapped, instance, args, kwargs):
|
|
941
|
+
"""Wrapper implementation for Vertex AI ImageGenerationModel.generate_images."""
|
|
942
|
+
logger.debug("Vertex AI ImageGenerationModel.generate_images wrapper called")
|
|
943
|
+
|
|
944
|
+
usage_metadata = getattr(instance, "_revenium_usage_metadata", {}) or kwargs.pop(
|
|
945
|
+
"usage_metadata", {}
|
|
946
|
+
)
|
|
947
|
+
|
|
948
|
+
model_name = _extract_model_name_from_instance(instance) or "imagen-3.0-generate-001"
|
|
949
|
+
|
|
950
|
+
# Extract image count from kwargs
|
|
951
|
+
number_of_images = kwargs.get("number_of_images", 1)
|
|
952
|
+
aspect_ratio = kwargs.get("aspect_ratio")
|
|
953
|
+
|
|
954
|
+
request_time_dt = datetime.datetime.now(datetime.timezone.utc)
|
|
955
|
+
|
|
956
|
+
# Call original
|
|
957
|
+
response = wrapped(*args, **kwargs)
|
|
958
|
+
|
|
959
|
+
response_time_dt = datetime.datetime.now(datetime.timezone.utc)
|
|
960
|
+
|
|
961
|
+
# Count generated images
|
|
962
|
+
actual_image_count = 0
|
|
963
|
+
if hasattr(response, "images") and response.images:
|
|
964
|
+
actual_image_count = len(response.images)
|
|
965
|
+
elif isinstance(response, list):
|
|
966
|
+
actual_image_count = len(response)
|
|
967
|
+
|
|
968
|
+
logger.debug(
|
|
969
|
+
f"Vertex AI Imagen generate_images: model={model_name}, "
|
|
970
|
+
f"requested={number_of_images}, actual={actual_image_count}"
|
|
971
|
+
)
|
|
972
|
+
|
|
973
|
+
try:
|
|
974
|
+
create_image_metering_call(
|
|
975
|
+
model=model_name,
|
|
976
|
+
requested_image_count=number_of_images,
|
|
977
|
+
actual_image_count=actual_image_count,
|
|
978
|
+
request_time_dt=request_time_dt,
|
|
979
|
+
response_time_dt=response_time_dt,
|
|
980
|
+
usage_metadata=usage_metadata,
|
|
981
|
+
operation_subtype="generation",
|
|
982
|
+
aspect_ratio=aspect_ratio,
|
|
983
|
+
)
|
|
984
|
+
except Exception as e:
|
|
985
|
+
logger.error(f"Error in Vertex AI Imagen metering: {e}")
|
|
986
|
+
|
|
987
|
+
return response
|
|
988
|
+
|
|
989
|
+
|
|
990
|
+
def _imagen_edit_image_impl(wrapped, instance, args, kwargs):
|
|
991
|
+
"""Wrapper implementation for Vertex AI ImageGenerationModel.edit_image."""
|
|
992
|
+
logger.debug("Vertex AI ImageGenerationModel.edit_image wrapper called")
|
|
993
|
+
|
|
994
|
+
usage_metadata = getattr(instance, "_revenium_usage_metadata", {}) or kwargs.pop(
|
|
995
|
+
"usage_metadata", {}
|
|
996
|
+
)
|
|
997
|
+
|
|
998
|
+
model_name = _extract_model_name_from_instance(instance) or "imagen-3.0-generate-001"
|
|
999
|
+
|
|
1000
|
+
number_of_images = kwargs.get("number_of_images", 1)
|
|
1001
|
+
|
|
1002
|
+
request_time_dt = datetime.datetime.now(datetime.timezone.utc)
|
|
1003
|
+
|
|
1004
|
+
response = wrapped(*args, **kwargs)
|
|
1005
|
+
|
|
1006
|
+
response_time_dt = datetime.datetime.now(datetime.timezone.utc)
|
|
1007
|
+
|
|
1008
|
+
actual_image_count = 0
|
|
1009
|
+
if hasattr(response, "images") and response.images:
|
|
1010
|
+
actual_image_count = len(response.images)
|
|
1011
|
+
|
|
1012
|
+
try:
|
|
1013
|
+
create_image_metering_call(
|
|
1014
|
+
model=model_name,
|
|
1015
|
+
requested_image_count=number_of_images,
|
|
1016
|
+
actual_image_count=actual_image_count,
|
|
1017
|
+
request_time_dt=request_time_dt,
|
|
1018
|
+
response_time_dt=response_time_dt,
|
|
1019
|
+
usage_metadata=usage_metadata,
|
|
1020
|
+
operation_subtype="edit",
|
|
1021
|
+
)
|
|
1022
|
+
except Exception as e:
|
|
1023
|
+
logger.error(f"Error in Vertex AI Imagen edit metering: {e}")
|
|
1024
|
+
|
|
1025
|
+
return response
|
|
1026
|
+
|
|
1027
|
+
|
|
1028
|
+
# --- Vertex AI Veo video generation wrapper ---
|
|
1029
|
+
|
|
1030
|
+
def _apply_veo_wrappers():
|
|
1031
|
+
"""
|
|
1032
|
+
Dynamically discover and wrap Vertex AI video generation models.
|
|
1033
|
+
Handles multiple module paths for forward compatibility.
|
|
1034
|
+
"""
|
|
1035
|
+
import importlib
|
|
1036
|
+
|
|
1037
|
+
module_patterns = [
|
|
1038
|
+
"vertexai.preview.vision_models",
|
|
1039
|
+
"vertexai.vision_models",
|
|
1040
|
+
]
|
|
1041
|
+
|
|
1042
|
+
wrapped_modules = []
|
|
1043
|
+
|
|
1044
|
+
for module_path in module_patterns:
|
|
1045
|
+
try:
|
|
1046
|
+
module = importlib.import_module(module_path)
|
|
1047
|
+
|
|
1048
|
+
# Check for VideoGenerationModel
|
|
1049
|
+
if hasattr(module, "VideoGenerationModel"):
|
|
1050
|
+
video_model_class = getattr(module, "VideoGenerationModel")
|
|
1051
|
+
|
|
1052
|
+
if hasattr(video_model_class, "generate_content"):
|
|
1053
|
+
@wrapt.patch_function_wrapper(
|
|
1054
|
+
module_path, "VideoGenerationModel.generate_content"
|
|
1055
|
+
)
|
|
1056
|
+
def veo_generate_wrapper(wrapped, instance, args, kwargs):
|
|
1057
|
+
return _veo_generate_impl(wrapped, instance, args, kwargs)
|
|
1058
|
+
|
|
1059
|
+
wrapped_modules.append(module_path + ".VideoGenerationModel.generate_content")
|
|
1060
|
+
|
|
1061
|
+
if hasattr(video_model_class, "generate"):
|
|
1062
|
+
@wrapt.patch_function_wrapper(
|
|
1063
|
+
module_path, "VideoGenerationModel.generate"
|
|
1064
|
+
)
|
|
1065
|
+
def veo_generate_alt_wrapper(wrapped, instance, args, kwargs):
|
|
1066
|
+
return _veo_generate_impl(wrapped, instance, args, kwargs)
|
|
1067
|
+
|
|
1068
|
+
wrapped_modules.append(module_path + ".VideoGenerationModel.generate")
|
|
1069
|
+
|
|
1070
|
+
except ImportError:
|
|
1071
|
+
logger.debug(f" Module {module_path} not available for Veo")
|
|
1072
|
+
except Exception as e:
|
|
1073
|
+
logger.debug(f" Error applying Veo wrapper to {module_path}: {e}")
|
|
1074
|
+
|
|
1075
|
+
if wrapped_modules:
|
|
1076
|
+
logger.info(f" Vertex AI Veo wrappers applied to: {', '.join(wrapped_modules)}")
|
|
1077
|
+
|
|
1078
|
+
return wrapped_modules
|
|
1079
|
+
|
|
1080
|
+
|
|
1081
|
+
def _veo_generate_impl(wrapped, instance, args, kwargs):
|
|
1082
|
+
"""Wrapper implementation for Vertex AI VideoGenerationModel.generate_content."""
|
|
1083
|
+
logger.debug("Vertex AI VideoGenerationModel wrapper called (Veo)")
|
|
1084
|
+
|
|
1085
|
+
usage_metadata = getattr(instance, "_revenium_usage_metadata", {}) or kwargs.pop(
|
|
1086
|
+
"usage_metadata", {}
|
|
1087
|
+
)
|
|
1088
|
+
|
|
1089
|
+
model_name = _extract_model_name_from_instance(instance) or "veo-2.0-generate-001"
|
|
1090
|
+
|
|
1091
|
+
# Extract video generation params
|
|
1092
|
+
duration = kwargs.get("duration", 5) # Default 5 seconds for Veo
|
|
1093
|
+
aspect_ratio = kwargs.get("aspect_ratio")
|
|
1094
|
+
|
|
1095
|
+
request_time_dt = datetime.datetime.now(datetime.timezone.utc)
|
|
1096
|
+
|
|
1097
|
+
# Call original
|
|
1098
|
+
response = wrapped(*args, **kwargs)
|
|
1099
|
+
|
|
1100
|
+
response_time_dt = datetime.datetime.now(datetime.timezone.utc)
|
|
1101
|
+
|
|
1102
|
+
# Extract video duration from response if available
|
|
1103
|
+
video_duration = float(duration)
|
|
1104
|
+
if hasattr(response, "duration_seconds"):
|
|
1105
|
+
video_duration = float(response.duration_seconds)
|
|
1106
|
+
elif hasattr(response, "duration"):
|
|
1107
|
+
video_duration = float(response.duration)
|
|
1108
|
+
|
|
1109
|
+
# Extract resolution if available
|
|
1110
|
+
resolution = None
|
|
1111
|
+
if hasattr(response, "resolution"):
|
|
1112
|
+
resolution = str(response.resolution)
|
|
1113
|
+
|
|
1114
|
+
# Extract video job ID for async operations
|
|
1115
|
+
video_job_id = None
|
|
1116
|
+
if hasattr(response, "operation_name"):
|
|
1117
|
+
video_job_id = str(response.operation_name)
|
|
1118
|
+
elif hasattr(response, "name"):
|
|
1119
|
+
video_job_id = str(response.name)
|
|
1120
|
+
|
|
1121
|
+
logger.debug(
|
|
1122
|
+
f"Vertex AI Veo generation: model={model_name}, "
|
|
1123
|
+
f"duration={video_duration}s, aspect_ratio={aspect_ratio}"
|
|
1124
|
+
)
|
|
1125
|
+
|
|
1126
|
+
# Auto-detect async operation when a job ID is present
|
|
1127
|
+
async_operation = video_job_id is not None
|
|
1128
|
+
|
|
1129
|
+
try:
|
|
1130
|
+
create_video_metering_call(
|
|
1131
|
+
model=model_name,
|
|
1132
|
+
duration_seconds=video_duration,
|
|
1133
|
+
request_time_dt=request_time_dt,
|
|
1134
|
+
response_time_dt=response_time_dt,
|
|
1135
|
+
usage_metadata=usage_metadata,
|
|
1136
|
+
operation_subtype="generation",
|
|
1137
|
+
resolution=resolution,
|
|
1138
|
+
aspect_ratio=aspect_ratio,
|
|
1139
|
+
video_job_id=video_job_id,
|
|
1140
|
+
async_operation=async_operation,
|
|
1141
|
+
)
|
|
1142
|
+
except Exception as e:
|
|
1143
|
+
logger.error(f"Error in Vertex AI Veo metering: {e}")
|
|
1144
|
+
|
|
1145
|
+
return response
|
|
1146
|
+
|
|
1147
|
+
|
|
1148
|
+
# Apply the dynamic wrappers when this module is imported
|
|
1149
|
+
try:
|
|
1150
|
+
_apply_generate_content_wrappers()
|
|
1151
|
+
except Exception as e:
|
|
1152
|
+
logger.error(f"Failed to apply dynamic Vertex AI wrappers: {e}")
|
|
1153
|
+
|
|
1154
|
+
try:
|
|
1155
|
+
_apply_imagen_wrappers()
|
|
1156
|
+
except Exception as e:
|
|
1157
|
+
logger.error(f"Failed to apply Vertex AI Imagen wrappers: {e}")
|
|
1158
|
+
|
|
1159
|
+
try:
|
|
1160
|
+
_apply_veo_wrappers()
|
|
1161
|
+
except Exception as e:
|
|
1162
|
+
logger.error(f"Failed to apply Vertex AI Veo wrappers: {e}")
|