From 072ddfdd0cc1cf9739816bf7a68afbce16b1c138 Mon Sep 17 00:00:00 2001 From: salmonumbrella <182032677+salmonumbrella@users.noreply.github.com> Date: Wed, 7 Oct 2026 08:16:41 +0000 Subject: [PATCH] feat(embedconfig): expose literal role affixes and dimension requests --- embedclient/config_contract_test.go | 160 ++++++++++++++++++++++++++++ embedconfig/embedder.go | 32 ++++-- embedconfig/embedder_test.go | 158 +++++++++++++++++++++++++++ embedconfig/example_test.go | 69 ++++++++++++ 4 files changed, 413 insertions(+), 6 deletions(-) create mode 100644 embedclient/config_contract_test.go create mode 100644 embedconfig/example_test.go diff --git a/embedclient/config_contract_test.go b/embedclient/config_contract_test.go new file mode 100644 index 0000000..410c9b6 --- /dev/null +++ b/embedclient/config_contract_test.go @@ -0,0 +1,160 @@ +package embedclient_test + +import ( + "encoding/base64" + "encoding/binary" + "fmt" + "math" + "net/http" + "net/http/httptest" + "testing" + + "github.com/BurntSushi/toml" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "go.kenn.io/kit/embedclient" + "go.kenn.io/kit/embedconfig" + "go.kenn.io/kit/embedmodel" +) + +// These tests exercise the text transport with synthetic vectors. They do +// not verify a model's tokenizer, pooling, checkpoint, or serving behavior. +func TestEmbedderTextContract(t *testing.T) { + for _, test := range []struct { + name string + dims int + requestDimensions bool + }{ + {"native", 768, false}, + {"requested native", 768, true}, + {"requested reduced", 512, true}, + {"server default reduced", 512, false}, + } { + t.Run(test.name, func(t *testing.T) { + requests := make(chan map[string]any, 1) + values := make([]float32, test.dims) + values[0], values[1] = 3, 4 + client := configuredTextClient(t, test.dims, test.requestDimensions, func(w http.ResponseWriter, r *http.Request) { + requests <- readBody(t, r) + writeJSON(t, w, map[string]any{ + "data": []map[string]any{{"index": 0, "embedding": values}}, + }) + }) + for _, input := range []struct { + role embedconfig.Role + prepared bool + text string + want string + }{ + {embedconfig.RoleDocument, false, "hello", "title: none | text: hello \n"}, + {embedconfig.RoleQuery, false, "hello", "task: search result | query: hello\t "}, + {embedconfig.RoleDocument, true, "title: none | text: hello \n", "title: none | text: hello \n"}, + {embedconfig.RoleQuery, true, "task: search result | query: hello\t ", "task: search result | query: hello\t "}, + } { + var vectors [][]float32 + var err error + if input.prepared { + vectors, err = client.EncodeFunc(input.role)(t.Context(), []string{input.text}) + } else { + vectors, err = client.EmbedTexts(t.Context(), input.role, []string{input.text}) + } + require.NoError(t, err) + require.Len(t, vectors, 1) + require.Len(t, vectors[0], test.dims) + assert.InDelta(t, 0.6, vectors[0][0], 1e-6) + assert.InDelta(t, 0.8, vectors[0][1], 1e-6) + var norm float64 + for _, value := range vectors[0] { + assert.False(t, math.IsNaN(float64(value)) || math.IsInf(float64(value), 0)) + norm += float64(value) * float64(value) + } + assert.InDelta(t, 1.0, norm, 1e-6) + body := <-requests + assert.Equal(t, "embed-text", body["model"]) + assert.Equal(t, []any{input.want}, body["input"]) + assert.NotContains(t, body, "input_type") + if test.requestDimensions { + assert.InDelta(t, float64(test.dims), body["dimensions"], 0) + } else { + assert.NotContains(t, body, "dimensions") + } + } + }) + } +} + +func TestEmbedderRejectsInvalidTextVectors(t *testing.T) { + nonFinite := make([]byte, 768*4) + binary.LittleEndian.PutUint32(nonFinite, math.Float32bits(float32(math.Inf(1)))) + native := make([]float32, 768) + native[0] = 1 + nullComponent := make([]any, 768) + for i := range nullComponent { + nullComponent[i] = 0 + } + nullComponent[0], nullComponent[1] = nil, 1 + for _, test := range []struct { + name string + dims int + requestDimensions bool + embedding any + }{ + {"wrong width", 768, false, []float32{3, 4}}, + {"native width instead of requested width", 512, true, native}, + {"zero norm", 768, false, make([]float32, 768)}, + {"null component", 768, false, nullComponent}, + {"nonfinite component", 768, false, base64.StdEncoding.EncodeToString(nonFinite)}, + } { + t.Run(test.name, func(t *testing.T) { + client := configuredTextClient(t, test.dims, test.requestDimensions, func(w http.ResponseWriter, _ *http.Request) { + writeJSON(t, w, map[string]any{ + "data": []map[string]any{{"index": 0, "embedding": test.embedding}}, + }) + }) + vectors, err := client.EmbedTexts(t.Context(), embedconfig.RoleDocument, []string{"hello"}) + require.ErrorIs(t, err, embedclient.ErrInvalidVector) + assert.Nil(t, vectors) + }) + } +} + +func TestEmbedderTextContractRejectsNonText(t *testing.T) { + client := configuredTextClient(t, 768, false, func(_ http.ResponseWriter, _ *http.Request) { + assert.Fail(t, "non-text input must be rejected before HTTP") + }) + for _, kind := range []string{embedmodel.KindImage, embedmodel.KindFile} { + vectors, err := client.Embed(t.Context(), []embedmodel.Content{{ + Role: embedconfig.RoleDocument, Kind: kind, Text: "hello", + }}) + require.ErrorIs(t, err, embedmodel.ErrUnsupportedContent) + assert.Nil(t, vectors) + } +} + +func configuredTextClient(t *testing.T, dims int, requestDimensions bool, handler http.HandlerFunc) *embedclient.Client { + t.Helper() + server := httptest.NewServer(handler) + t.Cleanup(server.Close) + var config embedconfig.Embedder + meta, err := toml.Decode(fmt.Sprintf(` +base_url = %q +model = "embed-text" +dims = %d +document_prefix = "title: none | text: " +document_suffix = " \n" +query_prefix = "task: search result | query: " +query_suffix = "\t " +request_dimensions = %t +`, server.URL, dims, requestDimensions), &config) + require.NoError(t, err) + require.Empty(t, meta.Undecoded()) + parts, err := config.Parts() + require.NoError(t, err) + client, err := embedclient.New(embedclient.Options{ + Model: parts.Model, Roles: parts.Roles, Deployment: parts.Deployment, + Batch: parts.Batch, Transport: parts.Transport, + }) + require.NoError(t, err) + return client +} diff --git a/embedconfig/embedder.go b/embedconfig/embedder.go index 875245b..8adbc43 100644 --- a/embedconfig/embedder.go +++ b/embedconfig/embedder.go @@ -23,6 +23,11 @@ type Embedder struct { Model string `toml:"model"` // Dims is the vector width the provider returns. Dims int `toml:"dims"` + // RequestDimensions sends Dims as the dimensions field on each request. + // Set it only when the endpoint supports requesting that width. False + // leaves the provider default; Dims still validates the returned width. + // The client never truncates vectors to fit Dims. + RequestDimensions bool `toml:"request_dimensions"` // APIKey is the bearer token: the token itself as a string, or a table // naming its source, such as { env = "NAME" } or { file = "PATH" }. // Leave it unset for an endpoint that needs no authentication. The @@ -35,6 +40,14 @@ type Embedder struct { // InputTypeMode is "none" (the default) or "retrieval". Retrieval sends // input_type document or query on each request. InputTypeMode string `toml:"input_type_mode"` + // DocumentPrefix, DocumentSuffix, QueryPrefix, and QuerySuffix are + // literal role affixes. Whitespace, including trailing spaces, is kept. + // They change vector identity and are applied to raw text by EmbedTexts + // and Embed; EncodeFunc expects text that is already formatted. + DocumentPrefix string `toml:"document_prefix"` + DocumentSuffix string `toml:"document_suffix"` + QueryPrefix string `toml:"query_prefix"` + QuerySuffix string `toml:"query_suffix"` // BatchSize caps inputs per request. Zero uses DefaultBatchItems. BatchSize int `toml:"batch_size"` // ModelContextTokens is the most tokens one input can hold, and @@ -91,16 +104,23 @@ func (e Embedder) Validate() error { // only operational defaults: batch size and timeout. func (e Embedder) Parts() (Parts, error) { model, err := Model{ - Name: e.Model, - Revision: e.FingerprintSalt, - Dimensions: e.Dims, - Metric: MetricCosine, - Normalization: NormalizationL2, + Name: e.Model, + Revision: e.FingerprintSalt, + Dimensions: e.Dims, + Metric: MetricCosine, + Normalization: NormalizationL2, + RequestDimensions: e.RequestDimensions, }.Prepared() if err != nil { return Parts{}, err } - roles, err := Roles{InputType: InputType(strings.TrimSpace(e.InputTypeMode))}.Prepared() + roles, err := Roles{ + InputType: InputType(strings.TrimSpace(e.InputTypeMode)), + DocumentPrefix: e.DocumentPrefix, + DocumentSuffix: e.DocumentSuffix, + QueryPrefix: e.QueryPrefix, + QuerySuffix: e.QuerySuffix, + }.Prepared() if err != nil { return Parts{}, err } diff --git a/embedconfig/embedder_test.go b/embedconfig/embedder_test.go index 7c31739..08604a9 100644 --- a/embedconfig/embedder_test.go +++ b/embedconfig/embedder_test.go @@ -9,6 +9,7 @@ import ( "github.com/stretchr/testify/require" "go.kenn.io/kit/embedconfig" + "go.kenn.io/kit/embedmodel" "go.kenn.io/kit/secretref" ) @@ -52,6 +53,163 @@ trust_private_network = true assert.Equal(t, 45*time.Second, parts.Transport.Timeout) } +func TestEmbedderLiteralRoleSettings(t *testing.T) { + var embedder embedconfig.Embedder + meta, err := toml.Decode(` +base_url = "https://api.example.test/v1" +model = "embed-text" +dims = 768 +document_prefix = "title: none | text: " +document_suffix = " \n" +query_prefix = "task: search result | query: " +query_suffix = "\t " +request_dimensions = true +`, &embedder) + require.NoError(t, err) + assert.Empty(t, meta.Undecoded()) + require.NoError(t, embedder.Validate()) + parts, err := embedder.Parts() + require.NoError(t, err) + assert.True(t, parts.Model.RequestDimensions) + assert.Equal(t, 768, parts.Model.Dimensions) + for _, test := range []struct { + role embedconfig.Role + want string + }{ + {embedconfig.RoleDocument, "title: none | text: hello \n"}, + {embedconfig.RoleQuery, "task: search result | query: hello\t "}, + } { + text, err := embedmodel.Format(test.role, "hello", parts.Roles) + require.NoError(t, err) + assert.Equal(t, test.want, text) + } +} + +func TestEmbedderRoleSettingsRequireEndpoint(t *testing.T) { + for _, setting := range []string{ + `document_prefix = " "`, `document_suffix = " "`, + `query_prefix = " "`, `query_suffix = " "`, `request_dimensions = true`, + } { + t.Run(setting, func(t *testing.T) { + var embedder embedconfig.Embedder + _, err := toml.Decode(setting, &embedder) + require.NoError(t, err) + assert.Error(t, embedder.Validate()) + }) + } +} + +func TestEmbedderPreservesDefaultIdentities(t *testing.T) { + const oldIdentity = "75da3207794697849064648e84a8e02dd2f873ea3690159bfc067972f0e1bb37" + const config = `base_url = "https://api.example.test/v1" +model = "embed-text" +dims = 768 +` + for _, settings := range []string{"", `document_prefix = "" +document_suffix = "" +query_prefix = "" +query_suffix = "" +request_dimensions = false +`} { + desc := embedderDescriptor(t, config+settings) + space, err := desc.VectorIdentity() + require.NoError(t, err) + assert.Equal(t, oldIdentity, space) + input, err := desc.InputIdentity() + require.NoError(t, err) + assert.Equal(t, oldIdentity, input) + gen, err := desc.Generation() + require.NoError(t, err) + assert.Equal(t, "ea14a3d46851bf40", gen.Fingerprint()) + assert.False(t, desc.Model.RequestDimensions) + } + for _, settings := range []string{ + `document_prefix = " "`, `document_suffix = " "`, + `query_prefix = " "`, `query_suffix = " "`, + `request_dimensions = true`, `fingerprint_salt = "weights-2"`, + } { + t.Run(settings, func(t *testing.T) { + desc := embedderDescriptor(t, config+settings) + space, err := desc.VectorIdentity() + require.NoError(t, err) + assert.NotEqual(t, oldIdentity, space) + input, err := desc.InputIdentity() + require.NoError(t, err) + assert.NotEqual(t, oldIdentity, input) + gen, err := desc.Generation() + require.NoError(t, err) + assert.NotEqual(t, "ea14a3d46851bf40", gen.Fingerprint()) + }) + } +} + +func embedderDescriptor(t *testing.T, config string) embedmodel.Descriptor { + t.Helper() + var embedder embedconfig.Embedder + _, err := toml.Decode(config, &embedder) + require.NoError(t, err) + parts, err := embedder.Parts() + require.NoError(t, err) + return embedmodel.Descriptor{Model: parts.Model, Roles: parts.Roles, Deployment: parts.Deployment} +} + +func FuzzEmbedderLiteralAffixes(f *testing.F) { + f.Add("title: none | text: ", "\n", "task: search result | query: ", " ", "hello") + f.Add("", "", "", "", "") + f.Add("\x00\n=\\", " \t", "é", "\xff", "文") + f.Fuzz(func(t *testing.T, documentPrefix, documentSuffix, queryPrefix, querySuffix, text string) { + parts, err := embedconfig.Embedder{ + BaseURL: "https://api.example.test/v1", Model: "embed-text", Dims: 768, + DocumentPrefix: documentPrefix, DocumentSuffix: documentSuffix, + QueryPrefix: queryPrefix, QuerySuffix: querySuffix, + }.Parts() + require.NoError(t, err) + for _, test := range []struct { + role embedconfig.Role + want string + }{ + {embedconfig.RoleDocument, documentPrefix + text + documentSuffix}, + {embedconfig.RoleQuery, queryPrefix + text + querySuffix}, + } { + got, err := embedmodel.Format(test.role, text, parts.Roles) + require.NoError(t, err) + assert.Equal(t, test.want, got) + } + }) +} + +func FuzzEmbedderAffixIdentities(f *testing.F) { + f.Add("", "", "", "", false) + f.Add("\n=\\", "\x00", "query: ", "\xff", true) + f.Fuzz(func(t *testing.T, documentPrefix, documentSuffix, queryPrefix, querySuffix string, requestDimensions bool) { + embedder := embedconfig.Embedder{ + BaseURL: "https://api.example.test/v1", Model: "embed-text", Dims: 768, + DocumentPrefix: documentPrefix, DocumentSuffix: documentSuffix, + QueryPrefix: queryPrefix, QuerySuffix: querySuffix, RequestDimensions: requestDimensions, + } + identity := func(e embedconfig.Embedder) string { + parts, err := e.Parts() + require.NoError(t, err) + got, err := (embedmodel.Descriptor{Model: parts.Model, Roles: parts.Roles}).VectorIdentity() + require.NoError(t, err) + return got + } + original := identity(embedder) + for _, change := range []func(*embedconfig.Embedder){ + func(e *embedconfig.Embedder) { e.DocumentPrefix += " " }, + func(e *embedconfig.Embedder) { e.DocumentSuffix += " " }, + func(e *embedconfig.Embedder) { e.QueryPrefix += " " }, + func(e *embedconfig.Embedder) { e.QuerySuffix += " " }, + func(e *embedconfig.Embedder) { e.RequestDimensions = !e.RequestDimensions }, + func(e *embedconfig.Embedder) { e.Dims = 512 }, + } { + changed := embedder + change(&changed) + assert.NotEqual(t, original, identity(changed)) + } + }) +} + func TestEmbedderPartsFillOnlyOperationalDefaults(t *testing.T) { parts, err := embedconfig.Embedder{ BaseURL: "http://127.0.0.1:11434/v1", Model: "nomic", Dims: 768, diff --git a/embedconfig/example_test.go b/embedconfig/example_test.go new file mode 100644 index 0000000..a804461 --- /dev/null +++ b/embedconfig/example_test.go @@ -0,0 +1,69 @@ +package embedconfig_test + +import ( + "fmt" + + "github.com/BurntSushi/toml" + + "go.kenn.io/kit/embedclient" + "go.kenn.io/kit/embedconfig" + "go.kenn.io/kit/embedmodel" +) + +// ExampleEmbedder_textRetrieval configures the text path for a caller-managed +// OpenAI-compatible deployment. The native width and retrieval recipe come +// from https://ai.google.dev/gemma/docs/embeddinggemma/model_card_2. +// This example constructs a client without making an inference request; +// transport tests use synthetic responses. Actual serving and prompt behavior +// remain untested here. +// +// The operator must bind the example alias to the chosen checkpoint, tokenizer, +// mean pooling including prompts, and bfloat16 or float32 activations. Neither +// the alias nor fingerprint_salt verifies server weights. The server must +// enforce the 8192-token limit including affixes and must not add prompts again. +// model_context_tokens is only a conservative batching bound, not token +// admission control. +// +// For a reduced width, set dims to the returned width and request_dimensions +// to true only if the endpoint supports that request and performs truncation +// with L2 renormalization. Kit validates width and normalizes accepted vectors; +// it never slices them. EmbedTexts applies affixes to raw text; EncodeFunc +// sends already formatted text unchanged. +func ExampleEmbedder_textRetrieval() { + var config embedconfig.Embedder + _, err := toml.Decode(` +base_url = "http://127.0.0.1:8080/v1" +model = "embeddinggemma-2-text-r914f7f8" +fingerprint_salt = "google/embeddinggemma-2@914f7f89142e33e77833254d9c9b90c3cef7303b" +dims = 768 +document_prefix = "title: none | text: " +query_prefix = "task: search result | query: " +request_dimensions = false +`, &config) + if err != nil { + panic(err) + } + parts, err := config.Parts() + if err != nil { + panic(err) + } + client, err := embedclient.New(embedclient.Options{ + Model: parts.Model, Roles: parts.Roles, Deployment: parts.Deployment, + Batch: parts.Batch, Transport: parts.Transport, + }) + if err != nil { + panic(err) + } + for _, role := range []embedconfig.Role{embedconfig.RoleDocument, embedconfig.RoleQuery} { + text, err := embedmodel.Format(role, "hello", parts.Roles) + if err != nil { + panic(err) + } + fmt.Println(text) + } + fmt.Println(parts.Model.Dimensions, parts.Model.RequestDimensions, client != nil) + // Output: + // title: none | text: hello + // task: search result | query: hello + // 768 false true +}