Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -263,9 +263,12 @@ def with_instrumentation(*args, **kwargs):
# the response stream is fully consumed (which is why
# ``start_as_current_span`` cannot be used directly).
kwargs = _apply_invoke_prompt_safety(span, kwargs, _BEDROCK_INVOKE_SPAN_NAME)
stream_start_time = time.time()
with trace.use_span(span, end_on_exit=False):
response = fn(*args, **kwargs)
_handle_stream_call(span, kwargs, response, metric_params, event_logger)
_handle_stream_call(
span, kwargs, response, metric_params, event_logger, stream_start_time
)

return response

Expand Down Expand Up @@ -303,18 +306,23 @@ def with_instrumentation(*args, **kwargs):
kwargs = _apply_converse_prompt_safety(span, kwargs, _BEDROCK_CONVERSE_SPAN_NAME)
# ST-10.4: see _instrumented_model_invoke_with_response_stream
# for the rationale on use_span(end_on_exit=False).
stream_start_time = time.time()
with trace.use_span(span, end_on_exit=False):
response = fn(*args, **kwargs)
if span.is_recording():
_handle_converse_stream(span, kwargs, response, metric_params, event_logger)
_handle_converse_stream(
span, kwargs, response, metric_params, event_logger, stream_start_time
)

return response

return with_instrumentation


@dont_throw
def _handle_stream_call(span, kwargs, response, metric_params, event_logger):
def _handle_stream_call(
span, kwargs, response, metric_params, event_logger, stream_start_time
):

(provider, model_vendor, model) = _get_vendor_model(kwargs.get("modelId"))
request_body = json.loads(kwargs.get("body"))
Expand Down Expand Up @@ -358,6 +366,7 @@ def stream_done(response_body):
StreamingWrapper(response["body"]),
span=span,
stream_done_callback=stream_done,
stream_start_time=stream_start_time,
)


Expand Down Expand Up @@ -427,7 +436,9 @@ def _handle_converse(span, kwargs, response, metric_params, event_logger):


@dont_throw
def _handle_converse_stream(span, kwargs, response, metric_params, event_logger):
def _handle_converse_stream(
span, kwargs, response, metric_params, event_logger, stream_start_time
):
(provider, model_vendor, model) = _get_vendor_model(kwargs.get("modelId"))

set_converse_model_span_attributes(span, provider, model, kwargs)
Expand All @@ -446,6 +457,7 @@ def _handle_converse_stream(span, kwargs, response, metric_params, event_logger)
model=model,
metric_params=metric_params,
event_logger=event_logger,
stream_start_time=stream_start_time,
)


Expand Down
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from __future__ import annotations

import json
import time

from wrapt import ObjectProxy

Expand All @@ -20,6 +21,10 @@
from opentelemetry.semconv_ai import LLMRequestTypeValues


FR_STREAMING_TIME_TO_FIRST_TOKEN_MS = "fortifyroot.llm.streaming.time_to_first_token_ms"
FR_STREAMING_TIME_TO_GENERATE_MS = "fortifyroot.llm.streaming.time_to_generate_ms"


class _BedrockChunkStreamingSafety:
def __init__(self, *, span, span_name: str, request_type: str):
self._streams = CompletionTextStreamGroup(
Expand Down Expand Up @@ -144,12 +149,17 @@ def _payload_text(self, payload):


class BedrockInvokeSafetyStreamingWrapper(ObjectProxy):
def __init__(self, response, *, span, stream_done_callback=None):
def __init__(self, response, *, span, stream_done_callback=None, stream_start_time=None):
super().__init__(response)

self._self_span = span
self._self_stream_done_callback = stream_done_callback
self._self_accumulating_body = {}
self._self_pending_event = None
self._self_stream_start_time = (
time.time() if stream_start_time is None else stream_start_time
)
self._self_first_token_time = None
self._self_streaming_safety = _BedrockChunkStreamingSafety(
span=span,
span_name=getattr(span, "name", "bedrock.completion"),
Expand All @@ -173,6 +183,7 @@ def __iter__(self):
self._accumulate_event(self._self_pending_event)
yield self._self_pending_event

self._finish_streaming_latency()
if self._self_stream_done_callback:
self._self_stream_done_callback(self._self_accumulating_body)

Expand All @@ -181,6 +192,8 @@ def _accumulate_event(self, event):
if not isinstance(payload, dict):
return

self._observe_streaming_text(self._self_streaming_safety._payload_text(payload))

event_type = payload.get("type")
if event_type is None:
self._accumulate_events(payload)
Expand Down Expand Up @@ -217,6 +230,25 @@ def _accumulate_events(self, payload):
else:
self._self_accumulating_body[key] = value

def _observe_streaming_text(self, text):
if not isinstance(text, str) or not text or self._self_first_token_time is not None:
return
now = time.time()
self._self_first_token_time = now
if self._self_span.is_recording():
self._self_span.set_attribute(
FR_STREAMING_TIME_TO_FIRST_TOKEN_MS,
round((now - self._self_stream_start_time) * 1000),
)

def _finish_streaming_latency(self):
if self._self_first_token_time is None or not self._self_span.is_recording():
return
self._self_span.set_attribute(
FR_STREAMING_TIME_TO_GENERATE_MS,
round((time.time() - self._self_first_token_time) * 1000),
)


class _BedrockConverseStreamingSafety:
def __init__(self, *, span, span_name: str):
Expand Down Expand Up @@ -321,6 +353,7 @@ def __init__(
model,
metric_params,
event_logger,
stream_start_time=None,
):
super().__init__(response)

Expand All @@ -333,6 +366,10 @@ def __init__(
self._self_role = "unknown"
self._self_response_msg = []
self._self_span_ended = False
self._self_stream_start_time = (
time.time() if stream_start_time is None else stream_start_time
)
self._self_first_token_time = None
self._self_streaming_safety = _BedrockConverseStreamingSafety(
span=span,
span_name=getattr(span, "name", "bedrock.converse"),
Expand All @@ -356,6 +393,7 @@ def __iter__(self):
yield self._self_pending_event

if not self._self_span_ended:
self._finish_streaming_latency()
self._self_span.end()
self._self_span_ended = True

Expand All @@ -369,10 +407,12 @@ def _observe_event(self, event):
if "contentBlockDelta" in event:
delta_text = ((event.get("contentBlockDelta") or {}).get("delta") or {}).get("text")
if isinstance(delta_text, str):
self._observe_streaming_text(delta_text)
self._self_response_msg.append(delta_text)
elif "contentBlockStart" in event:
start_text = ((event.get("contentBlockStart") or {}).get("start") or {}).get("text")
if isinstance(start_text, str):
self._observe_streaming_text(start_text)
self._self_response_msg.append(start_text)

if "messageStop" in event:
Expand Down Expand Up @@ -402,15 +442,38 @@ def _observe_event(self, event):
)
converse_usage_record(self._self_span, metadata, self._self_metric_params)
if not self._self_span_ended:
self._finish_streaming_latency()
self._self_span.end()
self._self_span_ended = True

def _observe_streaming_text(self, text):
if not isinstance(text, str) or not text or self._self_first_token_time is not None:
return
now = time.time()
self._self_first_token_time = now
if self._self_span.is_recording():
self._self_span.set_attribute(
FR_STREAMING_TIME_TO_FIRST_TOKEN_MS,
round((now - self._self_stream_start_time) * 1000),
)

def create_invoke_stream_wrapper(response, *, span, stream_done_callback=None):
def _finish_streaming_latency(self):
if self._self_first_token_time is None or not self._self_span.is_recording():
return
self._self_span.set_attribute(
FR_STREAMING_TIME_TO_GENERATE_MS,
round((time.time() - self._self_first_token_time) * 1000),
)


def create_invoke_stream_wrapper(
response, *, span, stream_done_callback=None, stream_start_time=None
):
return BedrockInvokeSafetyStreamingWrapper(
response,
span=span,
stream_done_callback=stream_done_callback,
stream_start_time=stream_start_time,
)


Expand All @@ -422,6 +485,7 @@ def create_converse_stream_wrapper(
model,
metric_params,
event_logger,
stream_start_time=None,
):
return BedrockConverseSafetyStream(
response,
Expand All @@ -430,6 +494,7 @@ def create_converse_stream_wrapper(
model=model,
metric_params=metric_params,
event_logger=event_logger,
stream_start_time=stream_start_time,
)


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,8 @@
)
from opentelemetry.instrumentation.bedrock.streaming_safety import (
BedrockConverseSafetyStream,
FR_STREAMING_TIME_TO_FIRST_TOKEN_MS,
FR_STREAMING_TIME_TO_GENERATE_MS,
_BedrockChunkStreamingSafety,
_BedrockConverseStreamingSafety,
_decode_chunk_event,
Expand Down Expand Up @@ -402,6 +404,52 @@ def test_invoke_streaming_wrapper_masks_chunk_bytes_and_accumulates_masked_body(
assert completed == [{"outputText": "masked-atail"}]


def test_invoke_streaming_wrapper_sets_latency_span_attrs(monkeypatch):
exporter, tracer = _test_span()
register_completion_safety_stream_factory(
lambda _: _FakeStreamSession(["a"], flush_result="")
)
ticks = iter([100.125, 101.250])
monkeypatch.setattr(
"opentelemetry.instrumentation.bedrock.streaming_safety.time.time",
lambda: next(ticks),
)

event = {
"chunk": {
"bytes": json.dumps(
{"contentBlockDelta": {"delta": {"text": "a"}}}
).encode("utf-8")
}
}
with tracer.start_as_current_span("bedrock.completion") as span:
wrapper = create_invoke_stream_wrapper(
[event],
span=span,
stream_start_time=100.0,
)
list(wrapper)

attrs = exporter.get_finished_spans()[0].attributes
assert attrs[FR_STREAMING_TIME_TO_FIRST_TOKEN_MS] == 125
assert attrs[FR_STREAMING_TIME_TO_GENERATE_MS] == 1125


def test_invoke_streaming_wrapper_omits_latency_attrs_for_empty_stream():
exporter, tracer = _test_span()
with tracer.start_as_current_span("bedrock.completion") as span:
wrapper = create_invoke_stream_wrapper(
[],
span=span,
stream_start_time=100.0,
)
list(wrapper)

attrs = exporter.get_finished_spans()[0].attributes
assert FR_STREAMING_TIME_TO_FIRST_TOKEN_MS not in attrs
assert FR_STREAMING_TIME_TO_GENERATE_MS not in attrs


def test_converse_streaming_wrapper_masks_deltas_before_span_attributes():
exporter, tracer = _test_span()
register_completion_safety_stream_factory(
Expand Down Expand Up @@ -441,6 +489,66 @@ def test_converse_streaming_wrapper_masks_deltas_before_span_attributes():
)


def test_converse_streaming_wrapper_sets_latency_span_attrs(monkeypatch):
exporter, tracer = _test_span()
register_completion_safety_stream_factory(
lambda _: _FakeStreamSession(["a"], flush_result="")
)
ticks = iter([200.050, 200.900])
monkeypatch.setattr(
"opentelemetry.instrumentation.bedrock.streaming_safety.time.time",
lambda: next(ticks),
)

events = [
{"messageStart": {"role": "assistant"}},
{"contentBlockDelta": {"delta": {"text": "a"}}},
{"metadata": {"usage": {"inputTokens": 1, "outputTokens": 1}}},
]
span = tracer.start_span("bedrock.converse")
wrapper = create_converse_stream_wrapper(
events,
span=span,
provider="AWS",
model="demo",
metric_params=SimpleNamespace(
guardrail_activation=SimpleNamespace(add=lambda *args, **kwargs: None),
token_histogram=None,
duration_histogram=None,
vendor="AWS",
model="demo",
is_stream=True,
start_time=0.0,
),
event_logger=None,
stream_start_time=200.0,
)
list(wrapper)

attrs = exporter.get_finished_spans()[0].attributes
assert attrs[FR_STREAMING_TIME_TO_FIRST_TOKEN_MS] == 50
assert attrs[FR_STREAMING_TIME_TO_GENERATE_MS] == 850


def test_converse_streaming_wrapper_omits_latency_attrs_for_empty_stream():
exporter, tracer = _test_span()
span = tracer.start_span("bedrock.converse")
wrapper = create_converse_stream_wrapper(
[],
span=span,
provider="AWS",
model="demo",
metric_params=SimpleNamespace(),
event_logger=None,
stream_start_time=200.0,
)
list(wrapper)

attrs = exporter.get_finished_spans()[0].attributes
assert FR_STREAMING_TIME_TO_FIRST_TOKEN_MS not in attrs
assert FR_STREAMING_TIME_TO_GENERATE_MS not in attrs


def test_bedrock_chunk_streaming_safety_covers_payload_variants_and_helpers():
_, tracer = _test_span()
register_completion_safety_stream_factory(
Expand Down
Loading