2026-07-01 06:53:30 +00:00
|
|
|
package chroma
|
|
|
|
|
|
|
|
|
|
import (
|
|
|
|
|
"context"
|
|
|
|
|
"encoding/json"
|
|
|
|
|
"fmt"
|
|
|
|
|
"io"
|
|
|
|
|
"net/http"
|
|
|
|
|
"strings"
|
|
|
|
|
"time"
|
|
|
|
|
|
|
|
|
|
"github.com/VictorVargas/rony-llm-agent/pkg/rag"
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
// Config holds the settings for the ChromaDB backend.
|
|
|
|
|
type Config struct {
|
|
|
|
|
BaseURL string // e.g. "http://localhost:8000"
|
|
|
|
|
Timeout int // request timeout in seconds (0 = default)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// Backend implements rag.Backend using ChromaDB's REST API.
|
|
|
|
|
type Backend struct {
|
|
|
|
|
baseURL string
|
|
|
|
|
http *http.Client
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// New creates a new ChromaDB backend.
|
|
|
|
|
func New(cfg Config) (*Backend, error) {
|
|
|
|
|
baseURL := cfg.BaseURL
|
|
|
|
|
if baseURL == "" {
|
|
|
|
|
baseURL = "http://localhost:8000"
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
timeout := time.Duration(cfg.Timeout) * time.Second
|
|
|
|
|
if timeout == 0 {
|
|
|
|
|
timeout = 30 * time.Second
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
return &Backend{
|
|
|
|
|
baseURL: baseURL,
|
|
|
|
|
http: &http.Client{
|
|
|
|
|
Timeout: timeout,
|
|
|
|
|
},
|
|
|
|
|
}, nil
|
|
|
|
|
}
|
|
|
|
|
|
2026-07-06 07:05:30 +00:00
|
|
|
func (b *Backend) Upsert(ctx context.Context, id string, vector []float32, content string, metadata map[string]string) error {
|
2026-07-01 06:53:30 +00:00
|
|
|
collection := "rony-memory"
|
|
|
|
|
|
|
|
|
|
embeddings := make([][]float64, 1)
|
|
|
|
|
for _, f := range vector {
|
|
|
|
|
embeddings[0] = append(embeddings[0], float64(f))
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
metadatas := make(map[string]interface{})
|
|
|
|
|
for k, v := range metadata {
|
|
|
|
|
metadatas[k] = v
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
reqBody, err := json.Marshal(map[string]interface{}{
|
2026-07-06 07:05:30 +00:00
|
|
|
"ids": []string{id},
|
2026-07-01 06:53:30 +00:00
|
|
|
"embeddings": embeddings,
|
2026-07-06 07:05:30 +00:00
|
|
|
"documents": []string{content},
|
|
|
|
|
"metadatas": []map[string]interface{}{metadatas},
|
2026-07-01 06:53:30 +00:00
|
|
|
})
|
|
|
|
|
if err != nil {
|
|
|
|
|
return fmt.Errorf("marshaling upsert request: %w", err)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
endpoint := fmt.Sprintf("/api/v1/collections/%s/upsert", collection)
|
|
|
|
|
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, b.baseURL+endpoint, strings.NewReader(string(reqBody)))
|
|
|
|
|
if err != nil {
|
|
|
|
|
return fmt.Errorf("creating request: %w", err)
|
|
|
|
|
}
|
|
|
|
|
httpReq.Header.Set("Content-Type", "application/json")
|
|
|
|
|
|
|
|
|
|
resp, err := b.http.Do(httpReq)
|
|
|
|
|
if err != nil {
|
|
|
|
|
return fmt.Errorf("request failed: %w", err)
|
|
|
|
|
}
|
|
|
|
|
defer resp.Body.Close()
|
|
|
|
|
|
|
|
|
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
|
|
|
|
body, _ := io.ReadAll(resp.Body)
|
|
|
|
|
return fmt.Errorf("API error %d: %s", resp.StatusCode, string(body))
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
return nil
|
|
|
|
|
}
|
|
|
|
|
|
2026-07-06 07:05:30 +00:00
|
|
|
func (b *Backend) Search(ctx context.Context, _ string, queryVector []float32, topK int) ([]rag.SearchResult, error) {
|
|
|
|
|
if len(queryVector) == 0 {
|
|
|
|
|
return nil, fmt.Errorf("chroma backend requires an embedding vector; it has no lexical fallback")
|
|
|
|
|
}
|
|
|
|
|
|
2026-07-01 06:53:30 +00:00
|
|
|
collection := "rony-memory"
|
|
|
|
|
|
|
|
|
|
query := make([]float64, len(queryVector))
|
|
|
|
|
for i, f := range queryVector {
|
|
|
|
|
query[i] = float64(f)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
reqBody, err := json.Marshal(map[string]interface{}{
|
|
|
|
|
"queries": []map[string]interface{}{
|
|
|
|
|
{
|
|
|
|
|
"vector": query,
|
|
|
|
|
"n_results": topK,
|
|
|
|
|
},
|
|
|
|
|
},
|
|
|
|
|
})
|
|
|
|
|
if err != nil {
|
|
|
|
|
return nil, fmt.Errorf("marshaling search request: %w", err)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
endpoint := fmt.Sprintf("/api/v1/collections/%s/query", collection)
|
|
|
|
|
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, b.baseURL+endpoint, strings.NewReader(string(reqBody)))
|
|
|
|
|
if err != nil {
|
|
|
|
|
return nil, fmt.Errorf("creating request: %w", err)
|
|
|
|
|
}
|
|
|
|
|
httpReq.Header.Set("Content-Type", "application/json")
|
|
|
|
|
|
|
|
|
|
resp, err := b.http.Do(httpReq)
|
|
|
|
|
if err != nil {
|
|
|
|
|
return nil, fmt.Errorf("request failed: %w", err)
|
|
|
|
|
}
|
|
|
|
|
defer resp.Body.Close()
|
|
|
|
|
|
|
|
|
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
|
|
|
|
body, _ := io.ReadAll(resp.Body)
|
|
|
|
|
return nil, fmt.Errorf("API error %d: %s", resp.StatusCode, string(body))
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
var apiResp chromaQueryResponse
|
|
|
|
|
if err := json.NewDecoder(resp.Body).Decode(&apiResp); err != nil {
|
|
|
|
|
return nil, fmt.Errorf("decoding response: %w", err)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
if len(apiResp.Results) == 0 || len(apiResp.Results[0].IDs) == 0 {
|
|
|
|
|
return []rag.SearchResult{}, nil
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
results := make([]rag.SearchResult, len(apiResp.Results[0].IDs[0]))
|
|
|
|
|
for i := range apiResp.Results[0].IDs[0] {
|
|
|
|
|
var metadata map[string]string
|
|
|
|
|
if len(apiResp.Results[0].Metadatas) > 0 && len(apiResp.Results[0].Metadatas[0]) > i {
|
|
|
|
|
metadata = stringifyMap(apiResp.Results[0].Metadatas[0][i])
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
results[i] = rag.SearchResult{
|
|
|
|
|
ID: apiResp.Results[0].IDs[0][i],
|
|
|
|
|
Content: apiResp.Results[0].Documents[0][i],
|
|
|
|
|
Score: float32(apiResp.Results[0].Distances[0][i]),
|
|
|
|
|
Metadata: metadata,
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
return results, nil
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func (b *Backend) Forget(ctx context.Context, id string) error {
|
|
|
|
|
collection := "rony-memory"
|
|
|
|
|
|
|
|
|
|
reqBody, err := json.Marshal(map[string]interface{}{
|
|
|
|
|
"ids": []string{id},
|
|
|
|
|
})
|
|
|
|
|
if err != nil {
|
|
|
|
|
return fmt.Errorf("marshaling delete request: %w", err)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
endpoint := fmt.Sprintf("/api/v1/collections/%s/delete", collection)
|
|
|
|
|
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, b.baseURL+endpoint, strings.NewReader(string(reqBody)))
|
|
|
|
|
if err != nil {
|
|
|
|
|
return fmt.Errorf("creating request: %w", err)
|
|
|
|
|
}
|
|
|
|
|
httpReq.Header.Set("Content-Type", "application/json")
|
|
|
|
|
|
|
|
|
|
resp, err := b.http.Do(httpReq)
|
|
|
|
|
if err != nil {
|
|
|
|
|
return fmt.Errorf("request failed: %w", err)
|
|
|
|
|
}
|
|
|
|
|
defer resp.Body.Close()
|
|
|
|
|
|
|
|
|
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
|
|
|
|
body, _ := io.ReadAll(resp.Body)
|
|
|
|
|
return fmt.Errorf("API error %d: %s", resp.StatusCode, string(body))
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
return nil
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func (b *Backend) ForgetAll(ctx context.Context) error {
|
|
|
|
|
collection := "rony-memory"
|
|
|
|
|
|
|
|
|
|
reqBody := []byte(`{"where": {}}`)
|
|
|
|
|
|
|
|
|
|
endpoint := fmt.Sprintf("/api/v1/collections/%s/delete", collection)
|
|
|
|
|
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, b.baseURL+endpoint, strings.NewReader(string(reqBody)))
|
|
|
|
|
if err != nil {
|
|
|
|
|
return fmt.Errorf("creating request: %w", err)
|
|
|
|
|
}
|
|
|
|
|
httpReq.Header.Set("Content-Type", "application/json")
|
|
|
|
|
|
|
|
|
|
resp, err := b.http.Do(httpReq)
|
|
|
|
|
if err != nil {
|
|
|
|
|
return fmt.Errorf("request failed: %w", err)
|
|
|
|
|
}
|
|
|
|
|
defer resp.Body.Close()
|
|
|
|
|
|
|
|
|
|
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
|
|
|
|
body, _ := io.ReadAll(resp.Body)
|
|
|
|
|
return fmt.Errorf("API error %d: %s", resp.StatusCode, string(body))
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
return nil
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func stringifyMap(m map[string]interface{}) map[string]string {
|
|
|
|
|
result := make(map[string]string)
|
|
|
|
|
for k, v := range m {
|
|
|
|
|
result[k] = fmt.Sprintf("%v", v)
|
|
|
|
|
}
|
|
|
|
|
return result
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// chromaQueryResponse represents the structure of a ChromaDB query response.
|
|
|
|
|
type chromaQueryResponse struct {
|
2026-07-12 23:14:15 +00:00
|
|
|
Names []string `json:"names"`
|
2026-07-01 06:53:30 +00:00
|
|
|
Results []chromaQueryResults `json:"results"`
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
type chromaQueryResults struct {
|
2026-07-12 23:14:15 +00:00
|
|
|
IDs [][]string `json:"ids"`
|
|
|
|
|
Documents [][]string `json:"documents"`
|
|
|
|
|
Distances [][]float64 `json:"distances"`
|
2026-07-01 06:53:30 +00:00
|
|
|
Metadatas [][]map[string]interface{} `json:"metadatas"`
|
|
|
|
|
}
|