Add Chutes embedding client with retry logic

This commit is contained in:
shenlan 2025-08-14 20:20:33 +08:00
parent 1d20e12de1
commit 73774ea633
8 changed files with 136 additions and 12 deletions

View File

@ -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)

View File

@ -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

View 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
}

View 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)
}
}

View 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
}

View File

@ -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
}

View File

@ -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)

View File

@ -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)