1- package memory
1+ package embedding
22
33import (
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
8686func 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
106106func 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
118118func 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
130130func 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
154154func 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
169169func 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
196196func 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
227227func 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
262262func 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