diff --git a/go.mod b/go.mod index 78c65dc..a1badec 100644 --- a/go.mod +++ b/go.mod @@ -3,7 +3,10 @@ module eko go 1.26 require ( + github.com/klauspost/compress v1.19.2 + github.com/lib/pq v1.12.3 github.com/mattn/go-sqlite3 v1.14.45 + github.com/pgvector/pgvector-go v0.4.1 github.com/spf13/cobra v1.10.2 go.opentelemetry.io/otel v1.45.0 go.opentelemetry.io/otel/exporters/otlp/otlpmetric/otlpmetrichttp v1.45.0 @@ -12,6 +15,7 @@ require ( go.opentelemetry.io/otel/sdk v1.45.0 go.opentelemetry.io/otel/sdk/metric v1.45.0 go.opentelemetry.io/otel/trace v1.45.0 + golang.org/x/sys v0.47.0 ) require ( @@ -22,17 +26,14 @@ require ( github.com/google/uuid v1.6.0 // indirect github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect - github.com/klauspost/compress v1.19.2 // indirect github.com/spf13/pflag v1.0.9 // indirect go.opentelemetry.io/auto/sdk v1.2.1 // indirect go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.45.0 // indirect go.opentelemetry.io/proto/otlp v1.11.0 // indirect golang.org/x/net v0.57.0 // indirect - golang.org/x/sys v0.47.0 // indirect golang.org/x/text v0.40.0 // indirect google.golang.org/genproto/googleapis/api v0.0.0-20260803160001-6ac0973c030d // indirect google.golang.org/genproto/googleapis/rpc v0.0.0-20260803160001-6ac0973c030d // indirect google.golang.org/grpc v1.83.0 // indirect google.golang.org/protobuf v1.36.11 // indirect - ) diff --git a/go.sum b/go.sum index 97baf7e..f5ed3dc 100644 --- a/go.sum +++ b/go.sum @@ -22,8 +22,12 @@ github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2 github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= github.com/klauspost/compress v1.19.2 h1:hMRETovs/pu/dVWN7zIT1PGG8t509MwT6bO7XSi26R8= github.com/klauspost/compress v1.19.2/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= +github.com/lib/pq v1.12.3 h1:tTWxr2YLKwIvK90ZXEw8GP7UFHtcbTtty8zsI+YjrfQ= +github.com/lib/pq v1.12.3/go.mod h1:/p+8NSbOcwzAEI7wiMXFlgydTwcgTr3OSKMsD2BitpA= github.com/mattn/go-sqlite3 v1.14.45 h1:6KA/spDguL3KV8rnybG7ezSaE4SeMR3KC9VbUoAQaIk= github.com/mattn/go-sqlite3 v1.14.45/go.mod h1:pjEuOr8IwzLJP2MfGeTb0A35jauH+C2kbHKBr7yXKVQ= +github.com/pgvector/pgvector-go v0.4.1 h1:Oaj0mC0Ky8KaTweNHHpLwyFlN6a0nUFoo1vgSFTEhPI= +github.com/pgvector/pgvector-go v0.4.1/go.mod h1:4fSXyjl1TYAIdByAql6JazKWRr2s7J0g4hcRY5cBFCk= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM= @@ -58,7 +62,6 @@ go.opentelemetry.io/proto/otlp v1.11.0/go.mod h1:SmVizdCOAm3XBtG1g1NnOdhW6jtddT7 go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= - golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE= golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU= golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= diff --git a/internal/db/pgvector.go b/internal/db/pgvector.go new file mode 100644 index 0000000..99fb73d --- /dev/null +++ b/internal/db/pgvector.go @@ -0,0 +1,163 @@ +package db + +import ( + "context" + "database/sql" + "fmt" + + _ "github.com/lib/pq" + pgvector "github.com/pgvector/pgvector-go" +) + +// PGVectorStore is a PostgreSQL + pgvector backed store for semantic search over +// code ASTs, diffs, and repository memory embeddings. +type PGVectorStore struct { + db *sql.DB +} + +// NewPGVectorStore opens a connection to the PostgreSQL database at dsn and +// ensures the pgvector extension and embeddings table exist. +// +// dsn format: "host= port=5432 user= password=

dbname= sslmode=disable" +func NewPGVectorStore(dsn string) (*PGVectorStore, error) { + db, err := sql.Open("postgres", dsn) + if err != nil { + return nil, fmt.Errorf("pgvector: open connection: %w", err) + } + if err := db.Ping(); err != nil { + return nil, fmt.Errorf("pgvector: ping: %w", err) + } + s := &PGVectorStore{db: db} + if err := s.migrate(); err != nil { + return nil, err + } + return s, nil +} + +// migrate creates the pgvector extension and the embeddings table if they don't +// already exist. The table uses an HNSW index for sub-millisecond ANN queries. +func (s *PGVectorStore) migrate() error { + stmts := []string{ + `CREATE EXTENSION IF NOT EXISTS vector`, + `CREATE TABLE IF NOT EXISTS embeddings ( + id TEXT PRIMARY KEY, + snapshot_id TEXT NOT NULL, + kind TEXT NOT NULL, + content TEXT NOT NULL, + embedding vector(1536) + )`, + // HNSW index for cosine similarity — best trade-off between recall and + // query latency for 1536-dim OpenAI-compatible embeddings. + `CREATE INDEX IF NOT EXISTS embeddings_hnsw_idx + ON embeddings USING hnsw (embedding vector_cosine_ops)`, + } + for _, stmt := range stmts { + if _, err := s.db.Exec(stmt); err != nil { + return fmt.Errorf("pgvector: migrate: %w", err) + } + } + return nil +} + +// Embedding holds a vector embedding for a piece of eko content. +type Embedding struct { + // ID is a stable content-derived identifier (e.g. SHA256 of content). + ID string + // SnapshotID is the eko snapshot this embedding belongs to. + SnapshotID string + // Kind classifies the content: "ast", "diff", or "memory". + Kind string + // Content is the raw text that was embedded. + Content string + // Vector is the 1536-dimensional float32 embedding. + Vector []float32 +} + +// Upsert stores an embedding, replacing any existing row with the same ID. +func (s *PGVectorStore) Upsert(ctx context.Context, e Embedding) error { + _, err := s.db.ExecContext(ctx, + `INSERT INTO embeddings (id, snapshot_id, kind, content, embedding) + VALUES ($1, $2, $3, $4, $5) + ON CONFLICT (id) DO UPDATE + SET snapshot_id = EXCLUDED.snapshot_id, + kind = EXCLUDED.kind, + content = EXCLUDED.content, + embedding = EXCLUDED.embedding`, + e.ID, e.SnapshotID, e.Kind, e.Content, + pgvector.NewVector(e.Vector), + ) + if err != nil { + return fmt.Errorf("pgvector: upsert %s: %w", e.ID, err) + } + return nil +} + +// SearchResult is a single result from a semantic similarity search. +type SearchResult struct { + Embedding + // Distance is the cosine distance from the query vector (lower = more similar). + Distance float64 +} + +// Search returns the top-k embeddings closest to query by cosine similarity. +// Pass kind="" to search across all content types, or "ast"/"diff"/"memory" to +// restrict the search. +func (s *PGVectorStore) Search(ctx context.Context, query []float32, kind string, topK int) ([]SearchResult, error) { + var ( + rows *sql.Rows + err error + ) + + qvec := pgvector.NewVector(query) + + if kind == "" { + rows, err = s.db.QueryContext(ctx, + `SELECT id, snapshot_id, kind, content, + embedding <=> $1 AS distance + FROM embeddings + ORDER BY distance + LIMIT $2`, + qvec, topK, + ) + } else { + rows, err = s.db.QueryContext(ctx, + `SELECT id, snapshot_id, kind, content, + embedding <=> $1 AS distance + FROM embeddings + WHERE kind = $2 + ORDER BY distance + LIMIT $3`, + qvec, kind, topK, + ) + } + if err != nil { + return nil, fmt.Errorf("pgvector: search: %w", err) + } + defer rows.Close() + + var results []SearchResult + for rows.Next() { + var r SearchResult + if err := rows.Scan(&r.ID, &r.SnapshotID, &r.Kind, &r.Content, &r.Distance); err != nil { + return nil, fmt.Errorf("pgvector: scan row: %w", err) + } + results = append(results, r) + } + return results, rows.Err() +} + +// DeleteBySnapshot removes all embeddings that belong to the given snapshot ID. +func (s *PGVectorStore) DeleteBySnapshot(ctx context.Context, snapshotID string) error { + _, err := s.db.ExecContext(ctx, + `DELETE FROM embeddings WHERE snapshot_id = $1`, snapshotID, + ) + if err != nil { + return fmt.Errorf("pgvector: delete snapshot %s: %w", snapshotID, err) + } + return nil +} + +// Close releases the database connection pool. +func (s *PGVectorStore) Close() error { + return s.db.Close() +}