Skip to content
Open
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
10 changes: 8 additions & 2 deletions ldai/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -114,7 +114,7 @@ func (c *Client) CreateTracker(token string, context ldcontext.Context) (*Tracke
// returns the resulting Config. Used for all error-path returns in evaluateConfig.
func (c *Client) returnDefault(key string, context ldcontext.Context, def Config) Config {
def.trackerFactory = func() *Tracker {
return newTracker(c.sdk, newRunID(), key, "", 1, context, &def, c.logger)
return newTracker(c.sdk, newRunID(), key, "", 1, "", 1, context, &def, c.logger)
}
return def
}
Expand Down Expand Up @@ -189,9 +189,15 @@ func (c *Client) evaluateConfig(
version = *parsed.Meta.Version
}

modelVersion := 1
if parsed.Meta.ModelVersion != nil {
modelVersion = *parsed.Meta.ModelVersion
}

variationKey := parsed.Meta.VariationKey
modelKey := parsed.Meta.ModelKey
cfg.trackerFactory = func() *Tracker {
return newTracker(c.sdk, newRunID(), key, variationKey, version, context, &cfg, c.logger)
return newTracker(c.sdk, newRunID(), key, variationKey, version, modelKey, modelVersion, context, &cfg, c.logger)
}

return cfg
Expand Down
78 changes: 78 additions & 0 deletions ldai/client_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -150,6 +150,82 @@ func TestParseModelName(t *testing.T) {
}
}

func TestParseModelKeyAndVersion(t *testing.T) {
// modelKey/modelVersion are intentionally not exposed on Config (they'd read as properties of
// the LLM itself, e.g. a version like "5.4"); the only place they surface is the tracker's
// stamped event data, mirroring variationKey/version.
tests := []struct {
name string
json []byte
expectedKey string
expectedVersion int
}{
{
name: "missing",
json: []byte(`{"model": {"name": "gpt-4"}}`),
expectedKey: "",
expectedVersion: 1,
},
{
name: "modelKey and modelVersion set",
json: []byte(`{"model": {"name": "gpt-4"}, "_ldMeta": {"modelKey": "my-model", "modelVersion": 2}}`),
expectedKey: "my-model",
expectedVersion: 2,
},
{
name: "modelVersion only",
json: []byte(`{"model": {"name": "gpt-4"}, "_ldMeta": {"modelVersion": 3}}`),
expectedKey: "",
expectedVersion: 3,
},
}

for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
mockSDK := newMockSDK(test.json, nil)
client, err := NewClient(mockSDK)
require.NoError(t, err)
require.NotNil(t, client)
mockSDK.events = nil

defaultVal := NewConfig().Enable().WithMessage("hello", datamodel.User).Build()
cfg := client.CompletionConfig("key", ldcontext.New("user"), defaultVal, nil)
tracker := cfg.CreateTracker()
require.NotNil(t, tracker)
assert.NoError(t, tracker.TrackSuccess())

require.NotEmpty(t, mockSDK.events)
data := mockSDK.events[len(mockSDK.events)-1].data
assert.Equal(t, test.expectedKey, data.GetByKey("modelKey").StringValue())
assert.Equal(t, test.expectedVersion, data.GetByKey("modelVersion").IntValue())
})
}
}

func TestCreateTrackerStampsModelKeyAndVersionOnTrackData(t *testing.T) {
configJSON := []byte(`{
"_ldMeta": {"variationKey": "var-1", "enabled": true, "version": 1, "modelKey": "my-model", "modelVersion": 2},
"model": {"name": "gpt-4"},
"provider": {"name": "openai"},
"messages": [{"content": "hello", "role": "user"}]
}`)

mockSDK := newMockSDK(configJSON, nil)
client, err := NewClient(mockSDK)
require.NoError(t, err)
mockSDK.events = nil

cfg := client.CompletionConfig("my-config", ldcontext.New("user"), Disabled(), nil)
tracker := cfg.CreateTracker()
require.NotNil(t, tracker)
assert.NoError(t, tracker.TrackSuccess())

require.NotEmpty(t, mockSDK.events)
data := mockSDK.events[len(mockSDK.events)-1].data
assert.Equal(t, "my-model", data.GetByKey("modelKey").StringValue())
assert.Equal(t, 2, data.GetByKey("modelVersion").IntValue())
}

func TestParseProviderName(t *testing.T) {
tests := []struct {
name string
Expand Down Expand Up @@ -1009,6 +1085,8 @@ func TestClient_CreateTracker_RoundTrip(t *testing.T) {
// modelName and providerName should be empty on reconstructed tracker
assert.Equal(t, "", feedbackEvent.data.GetByKey("modelName").StringValue())
assert.Equal(t, "", feedbackEvent.data.GetByKey("providerName").StringValue())
assert.False(t, feedbackEvent.data.GetByKey("modelKey").IsDefined())
assert.Equal(t, 1, feedbackEvent.data.GetByKey("modelVersion").IntValue())
}

func TestClient_CreateTracker_InvalidToken(t *testing.T) {
Expand Down
6 changes: 6 additions & 0 deletions ldai/datamodel/datamodel.go
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,12 @@ type Meta struct {

// Version is the version of the Variation.
Version *int `json:"version,omitempty"`

// ModelKey is the model's stable, unique key (distinct from Model.Name, which is not guaranteed unique).
ModelKey string `json:"modelKey,omitempty"`

// ModelVersion is the pinned version of the model that the variation references.
ModelVersion *int `json:"modelVersion,omitempty"`
}

// Model defines the serialization format for a model.
Expand Down
22 changes: 16 additions & 6 deletions ldai/tracker.go
Original file line number Diff line number Diff line change
Expand Up @@ -174,11 +174,14 @@ func newTracker(
key string,
variationKey string,
version int,
modelKey string,
modelVersion int,
ctx ldcontext.Context,
config *Config,
loggers interfaces.LDLoggers,
) *Tracker {
return newTrackerWithStopwatch(events, runID, key, variationKey, version, ctx, config, loggers, &defaultStopwatch{})
return newTrackerWithStopwatch(
events, runID, key, variationKey, version, modelKey, modelVersion, ctx, config, loggers, &defaultStopwatch{})
}

// newTrackerWithStopwatch creates a new Tracker with the specified runID, key, event sink, config, context, loggers,
Expand All @@ -189,6 +192,8 @@ func newTrackerWithStopwatch(
key string,
variationKey string,
version int,
modelKey string,
modelVersion int,
ctx ldcontext.Context,
config *Config,
loggers interfaces.LDLoggers,
Expand All @@ -203,7 +208,11 @@ func newTrackerWithStopwatch(
Set("configKey", ldvalue.String(key)).
Set("version", ldvalue.Int(version)).
Set("providerName", ldvalue.String(config.ProviderName())).
Set("modelName", ldvalue.String(config.ModelName()))
Set("modelName", ldvalue.String(config.ModelName())).
Set("modelVersion", ldvalue.Int(modelVersion))
if modelKey != "" {
builder.Set("modelKey", ldvalue.String(modelKey))
}
if variationKey != "" {
builder.Set("variationKey", ldvalue.String(variationKey))
}
Expand All @@ -230,7 +239,7 @@ func (t *Tracker) logWarning(format string, args ...interface{}) {

// ResumptionToken returns a URL-safe Base64-encoded token that can be used to reconstruct a tracker
// in a different process (e.g., for deferred feedback). The token contains the runId, configKey,
// variationKey, and version. It does not contain modelName or providerName.
// variationKey, and version. It does not contain modelName, providerName, modelKey, or modelVersion.
func (t *Tracker) ResumptionToken() string {
payload := resumptionPayload{
RunID: t.runID,
Expand All @@ -245,8 +254,8 @@ func (t *Tracker) ResumptionToken() string {
// TrackerFromResumptionToken reconstructs a Tracker from a resumption token and the given context.
// This is used for cross-process scenarios (e.g., deferred feedback) where the original tracker
// is no longer available but its runId must be reused. The token is obtained from Tracker.ResumptionToken().
// The reconstructed tracker will have empty modelName and providerName since these are not included
// in the token.
// The reconstructed tracker will have empty modelName, providerName, and modelKey, and modelVersion
// defaults to 1, since these are not included in the token.
func TrackerFromResumptionToken(token string, sdk ServerSDK, context ldcontext.Context) (*Tracker, error) {
decoded, err := base64.RawURLEncoding.DecodeString(token)
if err != nil {
Expand All @@ -262,7 +271,8 @@ func TrackerFromResumptionToken(token string, sdk ServerSDK, context ldcontext.C
Set("configKey", ldvalue.String(payload.ConfigKey)).
Set("version", ldvalue.Int(payload.Version)).
Set("providerName", ldvalue.String("")).
Set("modelName", ldvalue.String(""))
Set("modelName", ldvalue.String("")).
Set("modelVersion", ldvalue.Int(1))
if payload.VariationKey != "" {
builder.Set("variationKey", ldvalue.String(payload.VariationKey))
}
Expand Down
Loading
Loading