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
160 changes: 160 additions & 0 deletions embedclient/config_contract_test.go
Original file line number Diff line number Diff line change
@@ -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
}
32 changes: 26 additions & 6 deletions embedconfig/embedder.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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
}
Expand Down
158 changes: 158 additions & 0 deletions embedconfig/embedder_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)

Expand Down Expand Up @@ -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,
Expand Down
Loading