Add Chutes embedding client with retry logic
This commit is contained in:
parent
1d20e12de1
commit
73774ea633
@ -73,7 +73,7 @@ var rootCmd = &cobra.Command{
|
||||
case "ollama":
|
||||
embedder = embed.NewOllama(embCfg.Endpoint, embCfg.Model, embCfg.Dimension)
|
||||
case "chutes":
|
||||
embedder = embed.NewOpenAI(embCfg.Endpoint, embCfg.APIKey, "", embCfg.Dimension)
|
||||
embedder = embed.NewChutes(embCfg.Endpoint, embCfg.APIKey, embCfg.Dimension)
|
||||
default:
|
||||
if embCfg.Model != "" {
|
||||
embedder = embed.NewOpenAI(embCfg.Endpoint, embCfg.APIKey, embCfg.Model, embCfg.Dimension)
|
||||
|
||||
@ -30,7 +30,7 @@ provider:
|
||||
- 'moonshotai/Kimi-K2-Instruct'
|
||||
|
||||
embedding:
|
||||
endpoint: https://chutes-baai-bge-m3.chutes.ai/embed/v1/embeddings
|
||||
endpoint: https://chutes-baai-bge-m3.chutes.ai/embed
|
||||
token: "cpk_xxxxxxxxxxxxxxxxxx"
|
||||
dimension: 0 # 0 = 首次响应自动探测维度
|
||||
rate_limit_tpm: 120000
|
||||
|
||||
61
internal/rag/embed/chutes.go
Normal file
61
internal/rag/embed/chutes.go
Normal file
@ -0,0 +1,61 @@
|
||||
package embed
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Chutes implements the Embedder interface for Chutes embedding services.
|
||||
type Chutes struct {
|
||||
endpoint string
|
||||
token string
|
||||
dim int
|
||||
client *http.Client
|
||||
}
|
||||
|
||||
// NewChutes creates a new Chutes embedder.
|
||||
func NewChutes(endpoint, token string, dim int) *Chutes {
|
||||
return &Chutes{
|
||||
endpoint: endpoint,
|
||||
token: token,
|
||||
dim: dim,
|
||||
client: &http.Client{Timeout: 30 * time.Second},
|
||||
}
|
||||
}
|
||||
|
||||
// Dimension returns the embedding dimension if known.
|
||||
func (c *Chutes) Dimension() int { return c.dim }
|
||||
|
||||
// Embed posts texts to the Chutes embedding endpoint.
|
||||
func (c *Chutes) Embed(ctx context.Context, inputs []string) ([][]float32, int, error) {
|
||||
payload := map[string]any{"inputs": inputs}
|
||||
body, _ := json.Marshal(payload)
|
||||
headers := map[string]string{"Content-Type": "application/json"}
|
||||
if c.token != "" {
|
||||
headers["Authorization"] = "Bearer " + c.token
|
||||
}
|
||||
resp, err := postWithRetry(ctx, c.client, c.endpoint, body, headers)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode >= 300 {
|
||||
return nil, 0, &HTTPError{Code: resp.StatusCode, Status: fmt.Sprintf("embed failed: %s", resp.Status)}
|
||||
}
|
||||
var out struct {
|
||||
Data [][]float32 `json:"data"`
|
||||
}
|
||||
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
if len(out.Data) != len(inputs) {
|
||||
return nil, 0, fmt.Errorf("embedding count mismatch")
|
||||
}
|
||||
if c.dim == 0 && len(out.Data) > 0 {
|
||||
c.dim = len(out.Data[0])
|
||||
}
|
||||
return out.Data, 0, nil
|
||||
}
|
||||
25
internal/rag/embed/chutes_test.go
Normal file
25
internal/rag/embed/chutes_test.go
Normal file
@ -0,0 +1,25 @@
|
||||
package embed
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestChutesEmbed(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.Write([]byte(`{"data":[[0.1,0.2],[0.3,0.4]]}`))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
emb := NewChutes(srv.URL, "", 0)
|
||||
vecs, _, err := emb.Embed(context.Background(), []string{"a", "b"})
|
||||
if err != nil {
|
||||
t.Fatalf("Embed returned error: %v", err)
|
||||
}
|
||||
if len(vecs) != 2 || len(vecs[0]) != 2 {
|
||||
t.Fatalf("unexpected embedding: %#v", vecs)
|
||||
}
|
||||
}
|
||||
43
internal/rag/embed/http.go
Normal file
43
internal/rag/embed/http.go
Normal file
@ -0,0 +1,43 @@
|
||||
package embed
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"net/http"
|
||||
"time"
|
||||
)
|
||||
|
||||
// postWithRetry sends an HTTP POST request with basic retry and backoff for
|
||||
// rate limiting or transient server errors.
|
||||
func postWithRetry(ctx context.Context, client *http.Client, url string, body []byte, headers map[string]string) (*http.Response, error) {
|
||||
var resp *http.Response
|
||||
var err error
|
||||
for i := 0; i < 3; i++ {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for k, v := range headers {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
resp, err = client.Do(req)
|
||||
if err != nil {
|
||||
if i == 2 {
|
||||
return nil, err
|
||||
}
|
||||
} else if resp.StatusCode == http.StatusTooManyRequests || resp.StatusCode >= 500 {
|
||||
if i == 2 {
|
||||
return resp, nil
|
||||
}
|
||||
resp.Body.Close()
|
||||
} else {
|
||||
return resp, nil
|
||||
}
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
case <-time.After(time.Duration(1<<i) * time.Second):
|
||||
}
|
||||
}
|
||||
return resp, err
|
||||
}
|
||||
@ -1,7 +1,6 @@
|
||||
package embed
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
@ -40,15 +39,11 @@ func (o *OpenAI) Embed(ctx context.Context, inputs []string) ([][]float32, int,
|
||||
payload["model"] = o.model
|
||||
}
|
||||
b, _ := json.Marshal(payload)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, o.endpoint, bytes.NewReader(b))
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
headers := map[string]string{"Content-Type": "application/json"}
|
||||
if o.apiKey != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+o.apiKey)
|
||||
headers["Authorization"] = "Bearer " + o.apiKey
|
||||
}
|
||||
resp, err := o.client.Do(req)
|
||||
resp, err := postWithRetry(ctx, o.client, o.endpoint, b, headers)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
|
||||
@ -72,7 +72,7 @@ func IngestRepo(ctx context.Context, cfg *cfgpkg.Config, ds cfgpkg.DataSource, o
|
||||
case "ollama":
|
||||
embedder = embed.NewOllama(embCfg.Endpoint, embCfg.Model, embCfg.Dimension)
|
||||
case "chutes":
|
||||
embedder = embed.NewOpenAI(embCfg.Endpoint, embCfg.APIKey, "", embCfg.Dimension)
|
||||
embedder = embed.NewChutes(embCfg.Endpoint, embCfg.APIKey, embCfg.Dimension)
|
||||
default:
|
||||
if embCfg.Model != "" {
|
||||
embedder = embed.NewOpenAI(embCfg.Endpoint, embCfg.APIKey, embCfg.Model, embCfg.Dimension)
|
||||
|
||||
@ -67,7 +67,7 @@ func (s *Service) Query(ctx context.Context, question string, limit int) ([]Docu
|
||||
case "ollama":
|
||||
emb = embed.NewOllama(embCfg.Endpoint, embCfg.Model, embCfg.Dimension)
|
||||
case "chutes":
|
||||
emb = embed.NewOpenAI(embCfg.Endpoint, embCfg.APIKey, "", embCfg.Dimension)
|
||||
emb = embed.NewChutes(embCfg.Endpoint, embCfg.APIKey, embCfg.Dimension)
|
||||
default:
|
||||
if embCfg.Model != "" {
|
||||
emb = embed.NewOpenAI(embCfg.Endpoint, embCfg.APIKey, embCfg.Model, embCfg.Dimension)
|
||||
|
||||
Loading…
Reference in New Issue
Block a user