Skip to content

Commit 9d60500

Browse files
committed
Extract shared internal/embedding package from memory
Phase 1 of bringing sessions and skills to embedding parity with memory. The text-embedding backends (RandomProjections + OpenAI-compatible HTTP), their config shape, fingerprinting, and persistence were package-private inside internal/memory, so sessions and skills could not reuse them. Move them to a new internal/embedding package with an exported surface: - Config (was EmbeddingConfig) - TextEmbedder interface (exported methods) - New / NewRP constructors (was newTextEmbedder / newRPTextEmbedder) - Cosine similarity helper (was memory's cosineVector) - the bigram featurization (featurize.go) Memory keeps package-local aliases (EmbeddingConfig, textEmbedder) and thin constructors so all its call sites and the on-disk .gob/fingerprint formats are unchanged — no behavior change, no migration. cosineVector now delegates to embedding.Cosine to avoid a divergent copy. The embedder and featurize unit tests move with the code; memory keeps a local mock embeddings server for its episode-index integration tests.
1 parent f42f315 commit 9d60500

11 files changed

Lines changed: 481 additions & 361 deletions

File tree

Lines changed: 42 additions & 42 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
package memory
1+
package embedding
22

33
import (
44
"encoding/json"
@@ -75,16 +75,16 @@ func containsWord(text, w string) bool {
7575
return slices.Contains(strings.Fields(normalizeForEmbedding(text)), w)
7676
}
7777

78-
func httpCfg(srv *httptest.Server) *EmbeddingConfig {
79-
return &EmbeddingConfig{
78+
func httpCfg(srv *httptest.Server) *Config {
79+
return &Config{
8080
Provider: "http",
8181
BaseURL: srv.URL + "/v1",
8282
Model: "mock-embed",
8383
}
8484
}
8585

8686
func TestNewTextEmbedderDefaultsToRP(t *testing.T) {
87-
for _, cfg := range []*EmbeddingConfig{
87+
for _, cfg := range []*Config{
8888
nil,
8989
{},
9090
{Provider: "rp"},
@@ -93,31 +93,31 @@ func TestNewTextEmbedderDefaultsToRP(t *testing.T) {
9393
{Provider: "http", Model: "m"}, // missing base_url
9494
{Provider: "something-else", Model: "m"}, // unknown provider
9595
} {
96-
emb := newTextEmbedder(cfg, 64)
96+
emb := New(cfg, 64)
9797
if _, ok := emb.(*rpTextEmbedder); !ok {
98-
t.Errorf("newTextEmbedder(%+v) = %T, want *rpTextEmbedder", cfg, emb)
98+
t.Errorf("New(%+v) = %T, want *rpTextEmbedder", cfg, emb)
9999
}
100-
if got := emb.fingerprint(); got != "rp/64" {
100+
if got := emb.Fingerprint(); got != "rp/64" {
101101
t.Errorf("fingerprint = %q, want rp/64", got)
102102
}
103103
}
104104
}
105105

106106
func TestNewTextEmbedderHTTP(t *testing.T) {
107107
srv, _, _ := mockEmbedServer(t)
108-
emb := newTextEmbedder(httpCfg(srv), 64)
108+
emb := New(httpCfg(srv), 64)
109109
he, ok := emb.(*httpTextEmbedder)
110110
if !ok {
111-
t.Fatalf("newTextEmbedder = %T, want *httpTextEmbedder", emb)
111+
t.Fatalf("New = %T, want *httpTextEmbedder", emb)
112112
}
113-
if got := he.fingerprint(); got != "http/mock-embed/0" {
113+
if got := he.Fingerprint(); got != "http/mock-embed/0" {
114114
t.Errorf("fingerprint = %q, want http/mock-embed/0", got)
115115
}
116116
}
117117

118118
func TestNewTextEmbedderExpandsEnv(t *testing.T) {
119119
t.Setenv("ODEK_TEST_EMBED_URL", "http://localhost:9999/v1")
120-
emb := newTextEmbedder(&EmbeddingConfig{
120+
emb := New(&Config{
121121
Provider: "http",
122122
BaseURL: "${ODEK_TEST_EMBED_URL}",
123123
Model: "m",
@@ -129,36 +129,36 @@ func TestNewTextEmbedderExpandsEnv(t *testing.T) {
129129

130130
func TestHTTPEmbedderSemanticMatch(t *testing.T) {
131131
srv, _, _ := mockEmbedServer(t)
132-
emb := newTextEmbedder(httpCfg(srv), 64)
132+
emb := New(httpCfg(srv), 64)
133133

134-
a, err := emb.embed("the feline sat on the mat")
134+
a, err := emb.Embed("the feline sat on the mat")
135135
if err != nil {
136136
t.Fatal(err)
137137
}
138-
b, err := emb.embed("a cat appeared")
138+
b, err := emb.Embed("a cat appeared")
139139
if err != nil {
140140
t.Fatal(err)
141141
}
142-
c, err := emb.embed("postgres database migration")
142+
c, err := emb.Embed("postgres database migration")
143143
if err != nil {
144144
t.Fatal(err)
145145
}
146-
if simAB := cosineVector(a, b); simAB < 0.9 {
146+
if simAB := Cosine(a, b); simAB < 0.9 {
147147
t.Errorf("cat/feline cosine = %v, want ≥ 0.9 (semantic match)", simAB)
148148
}
149-
if simAC := cosineVector(a, c); simAC > 0.5 {
149+
if simAC := Cosine(a, c); simAC > 0.5 {
150150
t.Errorf("cat/database cosine = %v, want < 0.5", simAC)
151151
}
152152
}
153153

154154
func TestHTTPEmbedderCachesRepeatEmbeds(t *testing.T) {
155155
srv, requests, _ := mockEmbedServer(t)
156-
emb := newTextEmbedder(httpCfg(srv), 64)
156+
emb := New(httpCfg(srv), 64)
157157

158-
if _, err := emb.embed("hello world"); err != nil {
158+
if _, err := emb.Embed("hello world"); err != nil {
159159
t.Fatal(err)
160160
}
161-
if _, err := emb.embed("hello world"); err != nil {
161+
if _, err := emb.Embed("hello world"); err != nil {
162162
t.Fatal(err)
163163
}
164164
if got := requests.Load(); got != 1 {
@@ -168,10 +168,10 @@ func TestHTTPEmbedderCachesRepeatEmbeds(t *testing.T) {
168168

169169
func TestHTTPEmbedderFitBatchesOnlyMisses(t *testing.T) {
170170
srv, requests, texts := mockEmbedServer(t)
171-
emb := newTextEmbedder(httpCfg(srv), 64)
171+
emb := New(httpCfg(srv), 64)
172172

173173
corpus := []string{"one", "two", "three"}
174-
if err := emb.fit(corpus); err != nil {
174+
if err := emb.Fit(corpus); err != nil {
175175
t.Fatal(err)
176176
}
177177
if got := requests.Load(); got != 1 {
@@ -182,7 +182,7 @@ func TestHTTPEmbedderFitBatchesOnlyMisses(t *testing.T) {
182182
}
183183

184184
// Refit with one new entry: only the miss goes over the wire.
185-
if err := emb.fit(append(corpus, "four")); err != nil {
185+
if err := emb.Fit(append(corpus, "four")); err != nil {
186186
t.Fatal(err)
187187
}
188188
if got := requests.Load(); got != 2 {
@@ -195,9 +195,9 @@ func TestHTTPEmbedderFitBatchesOnlyMisses(t *testing.T) {
195195

196196
func TestHTTPEmbedderEmbedAllDedupsWithinBatch(t *testing.T) {
197197
srv, _, texts := mockEmbedServer(t)
198-
emb := newTextEmbedder(httpCfg(srv), 64)
198+
emb := New(httpCfg(srv), 64)
199199

200-
vecs, err := emb.embedAll([]string{"same", "same", "same"})
200+
vecs, err := emb.EmbedAll([]string{"same", "same", "same"})
201201
if err != nil {
202202
t.Fatal(err)
203203
}
@@ -214,62 +214,62 @@ func TestHTTPEmbedderErrorPropagates(t *testing.T) {
214214
http.Error(w, `{"error":{"message":"boom"}}`, http.StatusInternalServerError)
215215
}))
216216
defer srv.Close()
217-
emb := newTextEmbedder(&EmbeddingConfig{Provider: "http", BaseURL: srv.URL + "/v1", Model: "m"}, 64)
217+
emb := New(&Config{Provider: "http", BaseURL: srv.URL + "/v1", Model: "m"}, 64)
218218

219-
if _, err := emb.embed("x"); err == nil {
219+
if _, err := emb.Embed("x"); err == nil {
220220
t.Fatal("embed should propagate API errors")
221221
}
222-
if err := emb.fit([]string{"a", "b"}); err == nil {
222+
if err := emb.Fit([]string{"a", "b"}); err == nil {
223223
t.Fatal("fit should propagate API errors")
224224
}
225225
}
226226

227227
func TestRPTextEmbedderRoundTrip(t *testing.T) {
228-
emb := newRPTextEmbedder(64)
228+
emb := NewRP(64)
229229
corpus := []string{"uses postgres for storage", "prefers tabs over spaces"}
230-
if err := emb.fit(corpus); err != nil {
230+
if err := emb.Fit(corpus); err != nil {
231231
t.Fatal(err)
232232
}
233-
vecs, err := emb.embedAll(corpus)
233+
vecs, err := emb.EmbedAll(corpus)
234234
if err != nil {
235235
t.Fatal(err)
236236
}
237-
q, err := emb.embed("postgres storage")
237+
q, err := emb.Embed("postgres storage")
238238
if err != nil {
239239
t.Fatal(err)
240240
}
241-
if simSame := cosineVector(q, vecs[0]); simSame <= cosineVector(q, vecs[1]) {
241+
if simSame := Cosine(q, vecs[0]); simSame <= Cosine(q, vecs[1]) {
242242
t.Errorf("query should be closer to the postgres entry: %v vs %v",
243-
simSame, cosineVector(q, vecs[1]))
243+
simSame, Cosine(q, vecs[1]))
244244
}
245245

246246
// Persistence round-trip.
247247
path := t.TempDir() + "/rp.gob"
248-
emb.saveState(path)
249-
emb2 := newRPTextEmbedder(64)
250-
if !emb2.loadState(path) {
248+
emb.SaveState(path)
249+
emb2 := NewRP(64)
250+
if !emb2.LoadState(path) {
251251
t.Fatal("loadState failed")
252252
}
253-
q2, err := emb2.embed("postgres storage")
253+
q2, err := emb2.Embed("postgres storage")
254254
if err != nil {
255255
t.Fatal(err)
256256
}
257-
if cosineVector(q, q2) < 0.999 {
258-
t.Errorf("loaded embedder should reproduce vectors, cosine = %v", cosineVector(q, q2))
257+
if Cosine(q, q2) < 0.999 {
258+
t.Errorf("loaded embedder should reproduce vectors, cosine = %v", Cosine(q, q2))
259259
}
260260
}
261261

262262
func TestHTTPEmbedderCacheResetWhenFull(t *testing.T) {
263263
srv, _, _ := mockEmbedServer(t)
264-
emb := newTextEmbedder(httpCfg(srv), 64).(*httpTextEmbedder)
264+
emb := New(httpCfg(srv), 64).(*httpTextEmbedder)
265265

266266
// Fill past the cap in chunks; the cache must reset, not grow unbounded.
267267
batch := make([]string, 512)
268268
for round := range 10 {
269269
for i := range batch {
270270
batch[i] = fmt.Sprintf("text-%d-%d", round, i)
271271
}
272-
if _, err := emb.embedAll(batch); err != nil {
272+
if _, err := emb.EmbedAll(batch); err != nil {
273273
t.Fatal(err)
274274
}
275275
}

0 commit comments

Comments
 (0)