diff --git a/packages/sdk/server-ai/src/ldai/client.py b/packages/sdk/server-ai/src/ldai/client.py index 2d5ba42c..fd6993a6 100644 --- a/packages/sdk/server-ai/src/ldai/client.py +++ b/packages/sdk/server-ai/src/ldai/client.py @@ -1,4 +1,5 @@ import uuid +from dataclasses import dataclass from typing import Any, Callable, Dict, List, Optional, Tuple import chevron @@ -103,6 +104,34 @@ def _resolve_tools(variation: Dict[str, Any]) -> Optional[Dict[str, LDTool]]: return tools or None +@dataclass(frozen=True) +class _LdMeta: + """ + Parsed representation of a flag variation's ``_ldMeta`` block. + + Internal to :meth:`LDAIClient.__evaluate`. ``variation_key``, ``version``, + ``model_key``, and ``model_version`` are never exposed on a public config + type -- they only flow into the tracker built for that variation. + """ + enabled: bool + variation_key: str + version: int + model_key: Optional[str] + model_version: Optional[int] + + +def _parse_ld_meta(variation: Dict[str, Any]) -> _LdMeta: + ld_meta = variation.get('_ldMeta', {}) + raw_model_version = ld_meta.get('modelVersion') + return _LdMeta( + enabled=bool(ld_meta.get('enabled', False)), + variation_key=ld_meta.get('variationKey', ''), + version=int(ld_meta.get('version', 1)), + model_key=ld_meta.get('modelKey'), + model_version=int(raw_model_version) if raw_model_version is not None else None, + ) + + class LDAIClient: """The LaunchDarkly AI SDK client object.""" @@ -929,6 +958,8 @@ def __evaluate( provider = variation['provider'] provider_config = ProviderConfig(provider.get('name', '')) + meta = _parse_ld_meta(variation) + model = None if 'model' in variation and isinstance(variation['model'], dict): parameters = variation['model'].get('parameters', None) @@ -941,8 +972,6 @@ def __evaluate( region=region, ) - variation_key = variation.get('_ldMeta', {}).get('variationKey', '') - version = int(variation.get('_ldMeta', {}).get('version', 1)) model_name = model.name if model else '' provider_name = provider_config.name if provider_config else '' @@ -951,15 +980,17 @@ def tracker_factory() -> LDAIConfigTracker: ld_client=self._client, run_id=str(uuid.uuid4()), config_key=key, - variation_key=variation_key, - version=version, + variation_key=meta.variation_key, + version=meta.version, context=context, model_name=model_name, provider_name=provider_name, + model_key=meta.model_key, + model_version=meta.model_version, graph_key=graph_key, ) - enabled = variation.get('_ldMeta', {}).get('enabled', False) + enabled = meta.enabled judge_configuration = None if 'judgeConfiguration' in variation and isinstance(variation['judgeConfiguration'], dict): diff --git a/packages/sdk/server-ai/src/ldai/tracker.py b/packages/sdk/server-ai/src/ldai/tracker.py index 1cc5ee3e..2288c03b 100644 --- a/packages/sdk/server-ai/src/ldai/tracker.py +++ b/packages/sdk/server-ai/src/ldai/tracker.py @@ -110,6 +110,8 @@ def __init__( context: Context, model_name: str, provider_name: str, + model_key: Optional[str] = None, + model_version: Optional[int] = None, graph_key: Optional[str] = None, ): """ @@ -123,6 +125,8 @@ def __init__( :param context: Context for evaluation. :param model_name: Name of the model used. :param provider_name: Name of the provider used. + :param model_key: Stable, unique key of the model used. + :param model_version: Pinned version of the model used, when present in the payload. :param graph_key: When set, include ``graphKey`` in all event payloads (e.g. config-level metrics inside a graph). """ @@ -132,6 +136,8 @@ def __init__( self._version = version self._model_name = model_name self._provider_name = provider_name + self._model_key = model_key + self._model_version = model_version self._context = context self._graph_key = graph_key self._run_id = run_id @@ -147,7 +153,8 @@ def resumption_token(self) -> str: The token contains ``runId``, ``configKey``, ``version``, and optionally ``variationKey`` and ``graphKey`` (omitted when empty). - ``modelName`` and ``providerName`` are **not** included. + ``modelName``, ``providerName``, ``modelKey``, and ``modelVersion`` are + **not** included. """ data: dict = { "runId": self._run_id, @@ -219,6 +226,10 @@ def __get_track_data(self) -> dict: } if self._variation_key: data["variationKey"] = self._variation_key + if self._model_key: + data["modelKey"] = self._model_key + if self._model_version is not None: + data["modelVersion"] = self._model_version if self._graph_key: data['graphKey'] = self._graph_key return data diff --git a/packages/sdk/server-ai/tests/test_model_config.py b/packages/sdk/server-ai/tests/test_model_config.py index 6d0f0147..06ecf6b2 100644 --- a/packages/sdk/server-ai/tests/test_model_config.py +++ b/packages/sdk/server-ai/tests/test_model_config.py @@ -560,3 +560,82 @@ def test_create_tracker_each_call_has_different_run_id(): run_id_1 = success_calls[0].args[2]['runId'] run_id_2 = success_calls[1].args[2]['runId'] assert run_id_1 != run_id_2 + + +def test_create_tracker_stamps_model_key_and_version_on_track_data(): + from unittest.mock import Mock + + mock_client = Mock() + mock_client.variation.return_value = { + '_ldMeta': { + 'enabled': True, 'variationKey': 'var-abc', 'version': 7, + 'modelKey': 'my-model', 'modelVersion': 2, + }, + 'model': { + 'name': 'gpt-4', + }, + 'provider': {'name': 'openai'}, + 'messages': [] + } + + client = LDAIClient(mock_client) + context = Context.create('user-key') + + config = client.completion_config('my-config-key', context) + tracker = config.create_tracker() + tracker.track_success() + + success_calls = [ + c for c in mock_client.track.call_args_list + if c.args[0] == '$ld:ai:generation:success' + ] + assert len(success_calls) == 1 + track_data = success_calls[0].args[2] + assert track_data['modelKey'] == 'my-model' + assert track_data['modelVersion'] == 2 + + +@pytest.mark.parametrize( + 'ld_meta_overrides,expected_model_key,expected_model_version', + [ + pytest.param({}, None, None, id='omits_model_version_when_absent'), + pytest.param({'modelVersion': 3}, None, 3, id='omits_model_key_when_absent'), + ], +) +def test_create_tracker_model_key_and_version_defaults( + ld_meta_overrides, expected_model_key, expected_model_version, +): + from unittest.mock import Mock + + mock_client = Mock() + mock_client.variation.return_value = { + '_ldMeta': { + 'enabled': True, 'variationKey': 'var-abc', 'version': 7, + **ld_meta_overrides, + }, + 'model': {'name': 'gpt-4'}, + 'provider': {'name': 'openai'}, + 'messages': [] + } + + client = LDAIClient(mock_client) + context = Context.create('user-key') + + config = client.completion_config('my-config-key', context) + tracker = config.create_tracker() + tracker.track_success() + + success_calls = [ + c for c in mock_client.track.call_args_list + if c.args[0] == '$ld:ai:generation:success' + ] + assert len(success_calls) == 1 + track_data = success_calls[0].args[2] + if expected_model_key is None: + assert 'modelKey' not in track_data + else: + assert track_data['modelKey'] == expected_model_key + if expected_model_version is None: + assert 'modelVersion' not in track_data + else: + assert track_data['modelVersion'] == expected_model_version diff --git a/packages/sdk/server-ai/tests/test_tracker.py b/packages/sdk/server-ai/tests/test_tracker.py index ff6201f4..9fce0185 100644 --- a/packages/sdk/server-ai/tests/test_tracker.py +++ b/packages/sdk/server-ai/tests/test_tracker.py @@ -270,6 +270,36 @@ def _base_td() -> dict: } +def test_track_data_includes_model_key_when_set(client: LDClient): + context = Context.create("user-key") + tracker = LDAIConfigTracker( + ld_client=client, run_id="test-run-id", config_key="config-key", + variation_key="variation-key", version=3, model_name="fakeModel", + provider_name="fakeProvider", context=context, + model_key="my-model", model_version=2, + ) + tracker.track_success() + + track_data = client.track.call_args[0][2] # type: ignore + assert track_data["modelKey"] == "my-model" + assert track_data["modelVersion"] == 2 + + +def test_track_data_omits_model_key_when_empty(client: LDClient): + context = Context.create("user-key") + tracker = LDAIConfigTracker( + ld_client=client, run_id="test-run-id", config_key="config-key", + variation_key="variation-key", version=3, model_name="fakeModel", + provider_name="fakeProvider", context=context, + model_key="", model_version=3, + ) + tracker.track_success() + + track_data = client.track.call_args[0][2] # type: ignore + assert "modelKey" not in track_data + assert track_data["modelVersion"] == 3 + + def test_config_tracker_includes_graph_key_when_provided(client: LDClient): context = Context.create("user-key") tracker = LDAIConfigTracker( @@ -774,6 +804,7 @@ def test_resumption_token_round_trip(client: LDClient): ld_client=client, run_id="test-run-id", config_key="cfg-key", variation_key="var-key", version=5, model_name="gpt-4", provider_name="openai", context=context, + model_key="my-model", model_version=2, ) token = tracker.resumption_token @@ -788,6 +819,8 @@ def test_resumption_token_round_trip(client: LDClient): # modelName and providerName should NOT be in the token assert "modelName" not in decoded assert "providerName" not in decoded + assert "modelKey" not in decoded + assert "modelVersion" not in decoded def test_resumption_token_has_no_padding(client: LDClient): @@ -953,6 +986,8 @@ def test_client_create_tracker_from_resumption_token(): # modelName and providerName are empty when reconstructed from token assert track_data["modelName"] == "" assert track_data["providerName"] == "" + assert "modelVersion" not in track_data + assert "modelKey" not in track_data # Context should be the new one, not the original assert feedback_calls[0].args[1] == context