diff --git a/ldai/client.go b/ldai/client.go index 76eddc5e..c8d85f33 100644 --- a/ldai/client.go +++ b/ldai/client.go @@ -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 } @@ -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 diff --git a/ldai/client_test.go b/ldai/client_test.go index 01805f1e..04fc88d9 100644 --- a/ldai/client_test.go +++ b/ldai/client_test.go @@ -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 @@ -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) { diff --git a/ldai/datamodel/datamodel.go b/ldai/datamodel/datamodel.go index b5de1961..8a149f53 100644 --- a/ldai/datamodel/datamodel.go +++ b/ldai/datamodel/datamodel.go @@ -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. diff --git a/ldai/tracker.go b/ldai/tracker.go index 66473315..b951c480 100644 --- a/ldai/tracker.go +++ b/ldai/tracker.go @@ -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, @@ -189,6 +192,8 @@ func newTrackerWithStopwatch( key string, variationKey string, version int, + modelKey string, + modelVersion int, ctx ldcontext.Context, config *Config, loggers interfaces.LDLoggers, @@ -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)) } @@ -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, @@ -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 { @@ -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)) } diff --git a/ldai/tracker_test.go b/ldai/tracker_test.go index ef7eac15..33ef0bec 100644 --- a/ldai/tracker_test.go +++ b/ldai/tracker_test.go @@ -39,23 +39,33 @@ func (m *mockEvents) TrackMetric(eventName string, context ldcontext.Context, me func TestTracker_NewPanicsWithNilConfig(t *testing.T) { assert.Panics(t, func() { - newTracker(newMockEvents(), newRunID(), "key", "variationKey", 1, ldcontext.New("key"), nil, nil) + newTracker(newMockEvents(), newRunID(), "key", "variationKey", 1, "", 1, ldcontext.New("key"), nil, nil) }) } func TestTracker_NewDoesNotPanicWithConfig(t *testing.T) { assert.NotPanics(t, func() { - newTracker(newMockEvents(), newRunID(), "key", "variationKey", 1, ldcontext.New("key"), &Config{}, nil) + newTracker(newMockEvents(), newRunID(), "key", "variationKey", 1, "", 1, ldcontext.New("key"), &Config{}, nil) }) } func makeTrackData(configKey, variationKey string, version int, config *Config, runId string) ldvalue.Value { + return makeTrackDataWithModel(configKey, variationKey, version, "", 1, config, runId) +} + +func makeTrackDataWithModel( + configKey, variationKey string, version int, modelKey string, modelVersion int, config *Config, runId string, +) ldvalue.Value { builder := ldvalue.ObjectBuild(). Set("runId", ldvalue.String(runId)). Set("configKey", ldvalue.String(configKey)). 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)) } @@ -73,7 +83,7 @@ func extractRunId(t *testing.T, events *mockEvents) string { func TestTracker_TrackSuccess(t *testing.T) { events := newMockEvents() config := &Config{} - tracker := newTracker(events, newRunID(), "key", "variationKey", 1, ldcontext.New("key"), config, nil) + tracker := newTracker(events, newRunID(), "key", "variationKey", 1, "", 1, ldcontext.New("key"), config, nil) assert.NoError(t, tracker.TrackSuccess()) runId := extractRunId(t, events) @@ -92,7 +102,7 @@ func TestTracker_TrackSuccess(t *testing.T) { func TestTracker_TrackError(t *testing.T) { events := newMockEvents() config := &Config{} - tracker := newTracker(events, newRunID(), "key", "variationKey", 2, ldcontext.New("key"), config, nil) + tracker := newTracker(events, newRunID(), "key", "variationKey", 2, "", 1, ldcontext.New("key"), config, nil) assert.NoError(t, tracker.TrackError()) runId := extractRunId(t, events) @@ -111,7 +121,7 @@ func TestTracker_TrackError(t *testing.T) { func TestTracker_TrackRequest(t *testing.T) { events := newMockEvents() config := &Config{} - tracker := newTracker(events, newRunID(), "key", "variationKey", 3, ldcontext.New("key"), config, nil) + tracker := newTracker(events, newRunID(), "key", "variationKey", 3, "", 1, ldcontext.New("key"), config, nil) expectedResponse := ProviderResponse{ Usage: TokenUsage{ @@ -173,7 +183,7 @@ func TestTracker_TrackRequestReceivesConfig(t *testing.T) { Enable(). Build() - tracker := newTracker(events, newRunID(), "key", "variationKey", 4, ldcontext.New("key"), &expectedConfig, nil) + tracker := newTracker(events, newRunID(), "key", "variationKey", 4, "", 1, ldcontext.New("key"), &expectedConfig, nil) var gotConfig *Config _, _ = tracker.TrackRequest(func(c *Config) (ProviderResponse, error) { @@ -197,7 +207,7 @@ func TestTracker_LatencyMeasuredIfNotProvided(t *testing.T) { config := &Config{} tracker := newTrackerWithStopwatch( - events, newRunID(), "key", "variationKey", 5, ldcontext.New("key"), config, nil, mockStopwatch(42*time.Millisecond)) + events, newRunID(), "key", "variationKey", 5, "", 1, ldcontext.New("key"), config, nil, mockStopwatch(42*time.Millisecond)) expectedResponse := ProviderResponse{ Usage: TokenUsage{ @@ -221,7 +231,7 @@ func TestTracker_LatencyMeasuredIfNotProvided(t *testing.T) { func TestTracker_TrackDuration(t *testing.T) { events := newMockEvents() config := &Config{} - tracker := newTracker(events, newRunID(), "key", "variationKey", 6, ldcontext.New("key"), config, nil) + tracker := newTracker(events, newRunID(), "key", "variationKey", 6, "", 1, ldcontext.New("key"), config, nil) assert.NoError(t, tracker.TrackDuration(time.Millisecond*10)) @@ -240,7 +250,7 @@ func TestTracker_TrackFeedback(t *testing.T) { t.Run("positive feedback", func(t *testing.T) { events := newMockEvents() config := &Config{} - tracker := newTracker(events, newRunID(), "key", "variationKey", 7, ldcontext.New("key"), config, nil) + tracker := newTracker(events, newRunID(), "key", "variationKey", 7, "", 1, ldcontext.New("key"), config, nil) assert.NoError(t, tracker.TrackFeedback(FeedbackPositive)) @@ -258,7 +268,7 @@ func TestTracker_TrackFeedback(t *testing.T) { t.Run("negative feedback", func(t *testing.T) { events := newMockEvents() config := &Config{} - tracker := newTracker(events, newRunID(), "key", "variationKey", 7, ldcontext.New("key"), config, nil) + tracker := newTracker(events, newRunID(), "key", "variationKey", 7, "", 1, ldcontext.New("key"), config, nil) assert.NoError(t, tracker.TrackFeedback(FeedbackNegative)) @@ -276,7 +286,7 @@ func TestTracker_TrackFeedback(t *testing.T) { t.Run("invalid feedback returns error", func(t *testing.T) { events := newMockEvents() config := &Config{} - tracker := newTracker(events, newRunID(), "key", "variationKey", 7, ldcontext.New("key"), config, nil) + tracker := newTracker(events, newRunID(), "key", "variationKey", 7, "", 1, ldcontext.New("key"), config, nil) assert.Error(t, tracker.TrackFeedback("not a valid feedback value")) assert.Empty(t, events.events) @@ -287,7 +297,7 @@ func TestTracker_TrackTokens(t *testing.T) { t.Run("only one field set, only one event", func(t *testing.T) { events := newMockEvents() config := &Config{} - tracker := newTracker(events, newRunID(), "key", "variationKey", 8, ldcontext.New("key"), config, nil) + tracker := newTracker(events, newRunID(), "key", "variationKey", 8, "", 1, ldcontext.New("key"), config, nil) assert.NoError(t, tracker.TrackTokens(TokenUsage{ Total: 42, @@ -307,7 +317,7 @@ func TestTracker_TrackTokens(t *testing.T) { t.Run("all fields set, all events", func(t *testing.T) { events := newMockEvents() config := &Config{} - tracker := newTracker(events, newRunID(), "key", "variationKey", 9, ldcontext.New("key"), config, nil) + tracker := newTracker(events, newRunID(), "key", "variationKey", 9, "", 1, ldcontext.New("key"), config, nil) assert.NoError(t, tracker.TrackTokens(TokenUsage{ Total: 42, @@ -344,7 +354,7 @@ func TestTracker_TrackTokens(t *testing.T) { func TestTracker_GetSummary(t *testing.T) { t.Run("empty summary when nothing tracked", func(t *testing.T) { events := newMockEvents() - tracker := newTracker(events, newRunID(), "key", "variationKey", 10, ldcontext.New("key"), &Config{}, nil) + tracker := newTracker(events, newRunID(), "key", "variationKey", 10, "", 1, ldcontext.New("key"), &Config{}, nil) summary := tracker.GetSummary() @@ -357,7 +367,7 @@ func TestTracker_GetSummary(t *testing.T) { t.Run("first duration is returned", func(t *testing.T) { events := newMockEvents() - tracker := newTracker(events, newRunID(), "key", "variationKey", 11, ldcontext.New("key"), &Config{}, events.log.Loggers) + tracker := newTracker(events, newRunID(), "key", "variationKey", 11, "", 1, ldcontext.New("key"), &Config{}, events.log.Loggers) _ = tracker.TrackDuration(time.Millisecond * 10) _ = tracker.TrackDuration(time.Millisecond * 20) @@ -370,7 +380,7 @@ func TestTracker_GetSummary(t *testing.T) { t.Run("first feedback is returned", func(t *testing.T) { events := newMockEvents() - tracker := newTracker(events, newRunID(), "key", "variationKey", 12, ldcontext.New("key"), &Config{}, events.log.Loggers) + tracker := newTracker(events, newRunID(), "key", "variationKey", 12, "", 1, ldcontext.New("key"), &Config{}, events.log.Loggers) _ = tracker.TrackFeedback(FeedbackPositive) _ = tracker.TrackFeedback(FeedbackNegative) @@ -383,7 +393,7 @@ func TestTracker_GetSummary(t *testing.T) { t.Run("success status tracked correctly", func(t *testing.T) { events := newMockEvents() - tracker := newTracker(events, newRunID(), "key", "variationKey", 13, ldcontext.New("key"), &Config{}, nil) + tracker := newTracker(events, newRunID(), "key", "variationKey", 13, "", 1, ldcontext.New("key"), &Config{}, nil) _ = tracker.TrackSuccess() @@ -395,7 +405,7 @@ func TestTracker_GetSummary(t *testing.T) { t.Run("time to first token is returned", func(t *testing.T) { events := newMockEvents() - tracker := newTracker(events, newRunID(), "key", "variationKey", 14, ldcontext.New("key"), &Config{}, nil) + tracker := newTracker(events, newRunID(), "key", "variationKey", 14, "", 1, ldcontext.New("key"), &Config{}, nil) duration := time.Millisecond * 30 _ = tracker.TrackTimeToFirstToken(duration) @@ -408,7 +418,7 @@ func TestTracker_GetSummary(t *testing.T) { t.Run("token usage is returned", func(t *testing.T) { events := newMockEvents() - tracker := newTracker(events, newRunID(), "key", "variationKey", 15, ldcontext.New("key"), &Config{}, nil) + tracker := newTracker(events, newRunID(), "key", "variationKey", 15, "", 1, ldcontext.New("key"), &Config{}, nil) usage := TokenUsage{ Total: 100, @@ -427,7 +437,7 @@ func TestTracker_GetSummary(t *testing.T) { func TestTracker_RunIdPresentInTrackData(t *testing.T) { events := newMockEvents() config := &Config{} - tracker := newTracker(events, newRunID(), "key", "variationKey", 1, ldcontext.New("key"), config, nil) + tracker := newTracker(events, newRunID(), "key", "variationKey", 1, "", 1, ldcontext.New("key"), config, nil) _ = tracker.TrackSuccess() require.NotEmpty(t, events.events) @@ -440,7 +450,7 @@ func TestTracker_AtMostOnce(t *testing.T) { t.Run("TrackDuration only tracks once", func(t *testing.T) { events := newMockEvents() config := &Config{} - tracker := newTracker(events, newRunID(), "key", "variationKey", 1, ldcontext.New("key"), config, events.log.Loggers) + tracker := newTracker(events, newRunID(), "key", "variationKey", 1, "", 1, ldcontext.New("key"), config, events.log.Loggers) assert.NoError(t, tracker.TrackDuration(10*time.Millisecond)) assert.NoError(t, tracker.TrackDuration(20*time.Millisecond)) @@ -457,7 +467,7 @@ func TestTracker_AtMostOnce(t *testing.T) { t.Run("TrackTimeToFirstToken only tracks once", func(t *testing.T) { events := newMockEvents() config := &Config{} - tracker := newTracker(events, newRunID(), "key", "variationKey", 1, ldcontext.New("key"), config, events.log.Loggers) + tracker := newTracker(events, newRunID(), "key", "variationKey", 1, "", 1, ldcontext.New("key"), config, events.log.Loggers) assert.NoError(t, tracker.TrackTimeToFirstToken(10*time.Millisecond)) assert.NoError(t, tracker.TrackTimeToFirstToken(20*time.Millisecond)) @@ -474,7 +484,7 @@ func TestTracker_AtMostOnce(t *testing.T) { t.Run("TrackTokens only tracks once", func(t *testing.T) { events := newMockEvents() config := &Config{} - tracker := newTracker(events, newRunID(), "key", "variationKey", 1, ldcontext.New("key"), config, events.log.Loggers) + tracker := newTracker(events, newRunID(), "key", "variationKey", 1, "", 1, ldcontext.New("key"), config, events.log.Loggers) assert.NoError(t, tracker.TrackTokens(TokenUsage{Total: 10})) assert.NoError(t, tracker.TrackTokens(TokenUsage{Total: 20})) @@ -491,7 +501,7 @@ func TestTracker_AtMostOnce(t *testing.T) { t.Run("TrackFeedback only tracks once", func(t *testing.T) { events := newMockEvents() config := &Config{} - tracker := newTracker(events, newRunID(), "key", "variationKey", 1, ldcontext.New("key"), config, events.log.Loggers) + tracker := newTracker(events, newRunID(), "key", "variationKey", 1, "", 1, ldcontext.New("key"), config, events.log.Loggers) assert.NoError(t, tracker.TrackFeedback(FeedbackPositive)) assert.NoError(t, tracker.TrackFeedback(FeedbackNegative)) @@ -508,7 +518,7 @@ func TestTracker_AtMostOnce(t *testing.T) { t.Run("TrackSuccess only tracks once", func(t *testing.T) { events := newMockEvents() config := &Config{} - tracker := newTracker(events, newRunID(), "key", "variationKey", 1, ldcontext.New("key"), config, events.log.Loggers) + tracker := newTracker(events, newRunID(), "key", "variationKey", 1, "", 1, ldcontext.New("key"), config, events.log.Loggers) assert.NoError(t, tracker.TrackSuccess()) assert.NoError(t, tracker.TrackSuccess()) @@ -525,7 +535,7 @@ func TestTracker_AtMostOnce(t *testing.T) { t.Run("TrackError only tracks once", func(t *testing.T) { events := newMockEvents() config := &Config{} - tracker := newTracker(events, newRunID(), "key", "variationKey", 1, ldcontext.New("key"), config, events.log.Loggers) + tracker := newTracker(events, newRunID(), "key", "variationKey", 1, "", 1, ldcontext.New("key"), config, events.log.Loggers) assert.NoError(t, tracker.TrackError()) assert.NoError(t, tracker.TrackError()) @@ -542,7 +552,7 @@ func TestTracker_AtMostOnce(t *testing.T) { t.Run("TrackSuccess then TrackError only tracks success", func(t *testing.T) { events := newMockEvents() config := &Config{} - tracker := newTracker(events, newRunID(), "key", "variationKey", 1, ldcontext.New("key"), config, events.log.Loggers) + tracker := newTracker(events, newRunID(), "key", "variationKey", 1, "", 1, ldcontext.New("key"), config, events.log.Loggers) assert.NoError(t, tracker.TrackSuccess()) assert.NoError(t, tracker.TrackError()) @@ -556,7 +566,7 @@ func TestTracker_ResumptionToken(t *testing.T) { t.Run("produces valid base64url-encoded token", func(t *testing.T) { events := newMockEvents() config := &Config{} - tracker := newTracker(events, newRunID(), "my-config", "var-1", 3, ldcontext.New("key"), config, nil) + tracker := newTracker(events, newRunID(), "my-config", "var-1", 3, "", 1, ldcontext.New("key"), config, nil) token := tracker.ResumptionToken() assert.NotEmpty(t, token) @@ -579,10 +589,13 @@ func TestTracker_ResumptionToken(t *testing.T) { assert.Equal(t, 3, payload.Version) }) - t.Run("does not include modelName or providerName", func(t *testing.T) { + t.Run("does not include modelName, providerName, modelKey, or modelVersion", func(t *testing.T) { events := newMockEvents() - config := NewConfig().WithModelName("gpt-4").WithProviderName("openai").Build() - tracker := newTracker(events, newRunID(), "key", "var", 1, ldcontext.New("key"), &config, nil) + config := NewConfig(). + WithModelName("gpt-4"). + WithProviderName("openai"). + Build() + tracker := newTracker(events, newRunID(), "key", "var", 1, "my-model", 2, ldcontext.New("key"), &config, nil) token := tracker.ResumptionToken() decoded, err := base64.RawURLEncoding.DecodeString(token) @@ -593,7 +606,54 @@ func TestTracker_ResumptionToken(t *testing.T) { _, hasModel := raw["modelName"] _, hasProvider := raw["providerName"] + _, hasModelKey := raw["modelKey"] + _, hasModelVersion := raw["modelVersion"] assert.False(t, hasModel, "token should not contain modelName") assert.False(t, hasProvider, "token should not contain providerName") + assert.False(t, hasModelKey, "token should not contain modelKey") + assert.False(t, hasModelVersion, "token should not contain modelVersion") + }) +} + +func TestTracker_TrackDataIncludesModelKeyAndVersion(t *testing.T) { + t.Run("includes modelKey and modelVersion when set on config", func(t *testing.T) { + events := newMockEvents() + config := NewConfig().WithModelName("gpt-4").Build() + tracker := newTracker(events, newRunID(), "key", "var", 1, "my-model", 2, ldcontext.New("key"), &config, nil) + assert.NoError(t, tracker.TrackSuccess()) + + require.Len(t, events.events, 1) + data := events.events[0].data + assert.Equal(t, "my-model", data.GetByKey("modelKey").StringValue()) + assert.Equal(t, 2, data.GetByKey("modelVersion").IntValue()) + }) + + t.Run("omits modelKey when empty but still includes modelVersion", func(t *testing.T) { + events := newMockEvents() + config := NewConfig().WithModelName("gpt-4").Build() + tracker := newTracker(events, newRunID(), "key", "var", 1, "", 1, ldcontext.New("key"), &config, nil) + assert.NoError(t, tracker.TrackSuccess()) + + require.Len(t, events.events, 1) + data := events.events[0].data + assert.False(t, data.GetByKey("modelKey").IsDefined()) + assert.Equal(t, 1, data.GetByKey("modelVersion").IntValue()) }) } + +func TestTrackerFromResumptionToken_ModelKeyAndVersionDefaults(t *testing.T) { + mockSDK := newMockSDK(nil, nil) + config := NewConfig().WithModelName("gpt-4").Build() + tracker := newTracker( + mockSDK, newRunID(), "key", "var", 1, "my-model", 2, ldcontext.New("key"), &config, mockSDK.log.Loggers) + + token := tracker.ResumptionToken() + reconstructed, err := TrackerFromResumptionToken(token, mockSDK, ldcontext.New("key")) + require.NoError(t, err) + assert.NoError(t, reconstructed.TrackSuccess()) + + require.Len(t, mockSDK.events, 1) + data := mockSDK.events[0].data + assert.False(t, data.GetByKey("modelKey").IsDefined()) + assert.Equal(t, 1, data.GetByKey("modelVersion").IntValue()) +}