Skip to content
Draft
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 @@ -5,6 +5,7 @@
import time
import types
import warnings
from collections.abc import Mapping
from typing import Collection, Optional, Union

from opentelemetry import context as context_api
Expand Down Expand Up @@ -110,13 +111,68 @@ def _set_span_attribute(span, name, value):
return


def _set_api_attributes(span):
def _safe_getattr(value, name):
"""Read optional SDK attributes without breaking the instrumented call."""
try:
return getattr(value, name, None)
except Exception:
return None


def _get_url_from_credentials(credentials):
"""Return a URL from current credential objects or legacy mappings."""
try:
if isinstance(credentials, Mapping):
url = credentials.get("url")
else:
url = getattr(credentials, "url", None)
except Exception:
return None

return url if isinstance(url, str) and url else None


def _get_url_from_client(client):
"""Return the endpoint exposed by current or legacy Watsonx clients."""
for attribute in ("credentials", "wml_credentials"):
credentials = _safe_getattr(client, attribute)
if url := _get_url_from_credentials(credentials):
return url
return None


def _get_api_base_url(instance=None, kwargs=None):
"""Resolve the configured Watsonx endpoint when exposed by the SDK."""
kwargs = kwargs or {}

if url := _get_url_from_credentials(kwargs.get("credentials")):
return url

if url := _get_url_from_client(kwargs.get("api_client")):
return url

url = _safe_getattr(instance, "url")
if isinstance(url, str) and url:
return url

for attribute in ("credentials", "wml_credentials"):
credentials = _safe_getattr(instance, attribute)
if url := _get_url_from_credentials(credentials):
return url

if url := _get_url_from_client(_safe_getattr(instance, "_client")):
return url

return None


def _set_api_attributes(span, instance=None, kwargs=None):
if not span.is_recording():
return
_set_span_attribute(
span,
WatsonxSpanAttributes.WATSONX_API_BASE,
"https://us-south.ml.cloud.ibm.com",
_get_api_base_url(instance, kwargs),
)
_set_span_attribute(span, WatsonxSpanAttributes.WATSONX_API_TYPE, "watsonx.ai")
_set_span_attribute(span, WatsonxSpanAttributes.WATSONX_API_VERSION, "1.0")
Expand Down Expand Up @@ -488,8 +544,8 @@ def wrapper(wrapped, instance, args, kwargs):


@dont_throw
def _handle_input(span, event_logger, name, instance, response_counter, args, kwargs):
_set_api_attributes(span)
def _handle_input(span, event_logger, name, instance, args, kwargs):
_set_api_attributes(span, instance, kwargs)

if "generate" in name:
set_model_input_attributes(span, instance)
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,124 @@
from types import SimpleNamespace
from unittest.mock import Mock

import pytest
from opentelemetry.instrumentation.watsonx import (
WatsonxSpanAttributes,
_get_api_base_url,
_handle_input,
_set_api_attributes,
)


def test_get_api_base_url_from_current_sdk_credentials():
instance = SimpleNamespace(
_client=SimpleNamespace(credentials=SimpleNamespace(url="https://eu-de.ml.cloud.ibm.com"))
)

assert _get_api_base_url(instance) == "https://eu-de.ml.cloud.ibm.com"


def test_get_api_base_url_from_legacy_sdk_credentials():
instance = SimpleNamespace(_client=SimpleNamespace(wml_credentials={"url": "https://watsonx.example.internal"}))

assert _get_api_base_url(instance) == "https://watsonx.example.internal"


@pytest.mark.parametrize(
("kwargs", "expected"),
[
(
{"credentials": {"url": "https://jp-tok.ml.cloud.ibm.com"}},
"https://jp-tok.ml.cloud.ibm.com",
),
(
{"api_client": SimpleNamespace(credentials=SimpleNamespace(url="https://watsonx.custom"))},
"https://watsonx.custom",
),
],
)
def test_get_api_base_url_from_constructor_arguments(kwargs, expected):
assert _get_api_base_url(kwargs=kwargs) == expected


def test_get_api_base_url_returns_none_when_properties_raise():
class BrokenClient:
@property
def credentials(self):
raise RuntimeError("credentials unavailable")

@property
def wml_credentials(self):
raise RuntimeError("legacy credentials unavailable")

class BrokenInstance:
_client = BrokenClient()

@property
def url(self):
raise RuntimeError("URL unavailable")

@property
def credentials(self):
raise RuntimeError("credentials unavailable")

@property
def wml_credentials(self):
raise RuntimeError("legacy credentials unavailable")

assert _get_api_base_url(BrokenInstance()) is None


def test_get_api_base_url_returns_none_when_credentials_are_missing():
assert _get_api_base_url(SimpleNamespace()) is None


@pytest.mark.parametrize("url", [None, "", object()])
def test_get_api_base_url_returns_none_for_invalid_urls(url):
instance = SimpleNamespace(_client=SimpleNamespace(credentials=SimpleNamespace(url=url)))

assert _get_api_base_url(instance) is None


def test_set_api_attributes_uses_instance_endpoint():
span = Mock()
span.is_recording.return_value = True
instance = SimpleNamespace(
_client=SimpleNamespace(credentials=SimpleNamespace(url="https://eu-de.ml.cloud.ibm.com"))
)

_set_api_attributes(span, instance)

span.set_attribute.assert_any_call(
WatsonxSpanAttributes.WATSONX_API_BASE,
"https://eu-de.ml.cloud.ibm.com",
)


def test_set_api_attributes_omits_api_base_when_endpoint_is_unavailable():
span = Mock()
span.is_recording.return_value = True

_set_api_attributes(span, SimpleNamespace())

assert all(
call.args[0] != WatsonxSpanAttributes.WATSONX_API_BASE
for call in span.set_attribute.call_args_list
)


def test_handle_input_sets_api_base_without_swallowing_signature_error():
span = Mock()
span.is_recording.return_value = True
instance = SimpleNamespace(
model_id="ibm/granite",
params=None,
_client=SimpleNamespace(wml_credentials={"url": "https://watsonx.example.internal"}),
)

_handle_input(span, None, "watsonx.generate", instance, (), {})

span.set_attribute.assert_any_call(
WatsonxSpanAttributes.WATSONX_API_BASE,
"https://watsonx.example.internal",
)