Extract standalone services from cluster in own pagesUnverified
2cc88ae parent: 465d821 modified
docs/specs.md +5 -3 | @@ -994,17 +994,19 @@ glean/ | ||
| 994 | 994 | │ │ └── scraper.go # Full article content scraper |
| 995 | 995 | │ ├── metrics/ |
| 996 | 996 | │ │ └── metrics.go # Prometheus metrics definitions |
| 997 | +│ ├── ai/ | |
| 998 | +│ │ ├── embed.go # Embedder interface + OpenAI-compatible implementation + vector helpers | |
| 999 | +│ │ └── llm.go # TextModel interface + OpenAI-compatible LLM implementation | |
| 997 | 1000 | │ ├── cluster/ |
| 998 | 1001 | │ │ ├── jaccard.go # Jaccard similarity computation |
| 999 | -│ │ ├── embed.go # Embedder interface + OpenAI-compatible implementation | |
| 1000 | -│ │ ├── llm.go # LLM client for language detection and text tasks | |
| 1001 | 1002 | │ │ ├── article.go # Article + feed embedding computation, vec0 KNN content boost, language detection |
| 1002 | 1003 | │ │ ├── scoring.go # Feed + people + article recommendation queries (on-demand) |
| 1003 | 1004 | │ │ ├── social.go # Incremental follow-distance computation (1-3 hop, dirty-flag) |
| 1004 | -│ │ ├── dismiss.go # Dismiss + impression tracking | |
| 1005 | 1005 | │ │ ├── weights.go # Bandit-style signal weight auto-tuning |
| 1006 | 1006 | │ │ ├── diversity.go # Post-query domain/category diversity filtering |
| 1007 | 1007 | │ │ └── cron.go # Background recomputation scheduler |
| 1008 | +│ ├── feedback/ | |
| 1009 | +│ │ └── feedback.go # Dismiss + impression tracking service | |
| 1008 | 1010 | │ ├── server/ |
| 1009 | 1011 | │ │ ├── server.go # HTTP server, router setup |
| 1010 | 1012 | │ │ ├── auth_handler.go # OAuth login/callback/register |
| @@ -994,17 +994,19 @@ glean/ | |||
| 994 | │ │ └── scraper.go # Full article content scraper | 994 | │ │ └── scraper.go # Full article content scraper |
| 995 | │ ├── metrics/ | 995 | │ ├── metrics/ |
| 996 | │ │ └── metrics.go # Prometheus metrics definitions | 996 | │ │ └── metrics.go # Prometheus metrics definitions |
| 997 | +│ ├── ai/ | ||
| 998 | +│ │ ├── embed.go # Embedder interface + OpenAI-compatible implementation + vector helpers | ||
| 999 | +│ │ └── llm.go # TextModel interface + OpenAI-compatible LLM implementation | ||
| 997 | │ ├── cluster/ | 1000 | │ ├── cluster/ |
| 998 | │ │ ├── jaccard.go # Jaccard similarity computation | 1001 | │ │ ├── jaccard.go # Jaccard similarity computation |
| 999 | -│ │ ├── embed.go # Embedder interface + OpenAI-compatible implementation | ||
| 1000 | -│ │ ├── llm.go # LLM client for language detection and text tasks | ||
| 1001 | │ │ ├── article.go # Article + feed embedding computation, vec0 KNN content boost, language detection | 1002 | │ │ ├── article.go # Article + feed embedding computation, vec0 KNN content boost, language detection |
| 1002 | │ │ ├── scoring.go # Feed + people + article recommendation queries (on-demand) | 1003 | │ │ ├── scoring.go # Feed + people + article recommendation queries (on-demand) |
| 1003 | │ │ ├── social.go # Incremental follow-distance computation (1-3 hop, dirty-flag) | 1004 | │ │ ├── social.go # Incremental follow-distance computation (1-3 hop, dirty-flag) |
| 1004 | -│ │ ├── dismiss.go # Dismiss + impression tracking | ||
| 1005 | │ │ ├── weights.go # Bandit-style signal weight auto-tuning | 1005 | │ │ ├── weights.go # Bandit-style signal weight auto-tuning |
| 1006 | │ │ ├── diversity.go # Post-query domain/category diversity filtering | 1006 | │ │ ├── diversity.go # Post-query domain/category diversity filtering |
| 1007 | │ │ └── cron.go # Background recomputation scheduler | 1007 | │ │ └── cron.go # Background recomputation scheduler |
| 1008 | +│ ├── feedback/ | ||
| 1009 | +│ │ └── feedback.go # Dismiss + impression tracking service | ||
| 1008 | │ ├── server/ | 1010 | │ ├── server/ |
| 1009 | │ │ ├── server.go # HTTP server, router setup | 1011 | │ │ ├── server.go # HTTP server, router setup |
| 1010 | │ │ ├── auth_handler.go # OAuth login/callback/register | 1012 | │ │ ├── auth_handler.go # OAuth login/callback/register |
renamed
internal/ai/embed.go +10 -11 | similarity index 76% | ||
| rename from internal/cluster/embed.go | ||
| rename to internal/ai/embed.go | ||
| @@ -1,4 +1,4 @@ | ||
| 1 | -package cluster | |
| 1 | +package ai | |
| 2 | 2 | |
| 3 | 3 | import ( |
| 4 | 4 | "context" |
| @@ -9,26 +9,25 @@ import ( | ||
| 9 | 9 | "github.com/openai/openai-go/option" |
| 10 | 10 | ) |
| 11 | 11 | |
| 12 | -// Embedder generates vector embeddings for text inputs. | |
| 13 | 12 | type Embedder interface { |
| 14 | 13 | Embed(ctx context.Context, texts []string, instruction string) ([][]float32, error) |
| 15 | 14 | Dimension() int |
| 16 | 15 | } |
| 17 | 16 | |
| 18 | -type EmbedderClient struct { | |
| 17 | +type embedder struct { | |
| 19 | 18 | client openai.Client |
| 20 | 19 | model string |
| 21 | 20 | dimension int |
| 22 | 21 | } |
| 23 | 22 | |
| 24 | -type EmbedderClientConfig struct { | |
| 23 | +type EmbedConfig struct { | |
| 25 | 24 | BaseURL string |
| 26 | 25 | APIKey string |
| 27 | 26 | Model string |
| 28 | 27 | Dimension int |
| 29 | 28 | } |
| 30 | 29 | |
| 31 | -func NewEmbedderClient(cfg EmbedderClientConfig) *EmbedderClient { | |
| 30 | +func NewEmbedder(cfg EmbedConfig) Embedder { | |
| 32 | 31 | opts := []option.RequestOption{} |
| 33 | 32 | if cfg.BaseURL != "" { |
| 34 | 33 | opts = append(opts, option.WithBaseURL(cfg.BaseURL)) |
| @@ -36,18 +35,18 @@ func NewEmbedderClient(cfg EmbedderClientConfig) *EmbedderClient { | ||
| 36 | 35 | if cfg.APIKey != "" { |
| 37 | 36 | opts = append(opts, option.WithAPIKey(cfg.APIKey)) |
| 38 | 37 | } |
| 39 | - return &EmbedderClient{ | |
| 38 | + return &embedder{ | |
| 40 | 39 | client: openai.NewClient(opts...), |
| 41 | 40 | model: cfg.Model, |
| 42 | 41 | dimension: cfg.Dimension, |
| 43 | 42 | } |
| 44 | 43 | } |
| 45 | 44 | |
| 46 | -func (e *EmbedderClient) Dimension() int { | |
| 45 | +func (e *embedder) Dimension() int { | |
| 47 | 46 | return e.dimension |
| 48 | 47 | } |
| 49 | 48 | |
| 50 | -func (e *EmbedderClient) Embed(ctx context.Context, texts []string, instruction string) ([][]float32, error) { | |
| 49 | +func (e *embedder) Embed(ctx context.Context, texts []string, instruction string) ([][]float32, error) { | |
| 51 | 50 | inputs := texts |
| 52 | 51 | if instruction != "" { |
| 53 | 52 | inputs = make([]string, len(texts)) |
| @@ -75,11 +74,11 @@ func (e *EmbedderClient) Embed(ctx context.Context, texts []string, instruction | ||
| 75 | 74 | return embeddings, nil |
| 76 | 75 | } |
| 77 | 76 | |
| 78 | -func avgEmbeddings(blobs [][]byte, dim int) ([]byte, error) { | |
| 77 | +func AvgEmbeddings(blobs [][]byte, dim int) ([]byte, error) { | |
| 79 | 78 | sum := make([]float32, dim) |
| 80 | 79 | count := 0 |
| 81 | 80 | for _, blob := range blobs { |
| 82 | - v := bytesToFloat32s(blob, dim) | |
| 81 | + v := BytesToFloat32s(blob, dim) | |
| 83 | 82 | if v == nil { |
| 84 | 83 | continue |
| 85 | 84 | } |
| @@ -97,7 +96,7 @@ func avgEmbeddings(blobs [][]byte, dim int) ([]byte, error) { | ||
| 97 | 96 | return vec.SerializeFloat32(sum) |
| 98 | 97 | } |
| 99 | 98 | |
| 100 | -func bytesToFloat32s(data []byte, expectedDim int) []float32 { | |
| 99 | +func BytesToFloat32s(data []byte, expectedDim int) []float32 { | |
| 101 | 100 | if len(data) != expectedDim*4 { |
| 102 | 101 | return nil |
| 103 | 102 | } |
| similarity index 76% | |||
| rename from internal/cluster/embed.go | |||
| rename to internal/ai/embed.go | |||
| @@ -1,4 +1,4 @@ | |||
| 1 | -package cluster | 1 | +package ai |
| 2 | 2 | ||
| 3 | import ( | 3 | import ( |
| 4 | "context" | 4 | "context" |
| @@ -9,26 +9,25 @@ import ( | |||
| 9 | "github.com/openai/openai-go/option" | 9 | "github.com/openai/openai-go/option" |
| 10 | ) | 10 | ) |
| 11 | 11 | ||
| 12 | -// Embedder generates vector embeddings for text inputs. | ||
| 13 | type Embedder interface { | 12 | type Embedder interface { |
| 14 | Embed(ctx context.Context, texts []string, instruction string) ([][]float32, error) | 13 | Embed(ctx context.Context, texts []string, instruction string) ([][]float32, error) |
| 15 | Dimension() int | 14 | Dimension() int |
| 16 | } | 15 | } |
| 17 | 16 | ||
| 18 | -type EmbedderClient struct { | 17 | +type embedder struct { |
| 19 | client openai.Client | 18 | client openai.Client |
| 20 | model string | 19 | model string |
| 21 | dimension int | 20 | dimension int |
| 22 | } | 21 | } |
| 23 | 22 | ||
| 24 | -type EmbedderClientConfig struct { | 23 | +type EmbedConfig struct { |
| 25 | BaseURL string | 24 | BaseURL string |
| 26 | APIKey string | 25 | APIKey string |
| 27 | Model string | 26 | Model string |
| 28 | Dimension int | 27 | Dimension int |
| 29 | } | 28 | } |
| 30 | 29 | ||
| 31 | -func NewEmbedderClient(cfg EmbedderClientConfig) *EmbedderClient { | 30 | +func NewEmbedder(cfg EmbedConfig) Embedder { |
| 32 | opts := []option.RequestOption{} | 31 | opts := []option.RequestOption{} |
| 33 | if cfg.BaseURL != "" { | 32 | if cfg.BaseURL != "" { |
| 34 | opts = append(opts, option.WithBaseURL(cfg.BaseURL)) | 33 | opts = append(opts, option.WithBaseURL(cfg.BaseURL)) |
| @@ -36,18 +35,18 @@ func NewEmbedderClient(cfg EmbedderClientConfig) *EmbedderClient { | |||
| 36 | if cfg.APIKey != "" { | 35 | if cfg.APIKey != "" { |
| 37 | opts = append(opts, option.WithAPIKey(cfg.APIKey)) | 36 | opts = append(opts, option.WithAPIKey(cfg.APIKey)) |
| 38 | } | 37 | } |
| 39 | - return &EmbedderClient{ | 38 | + return &embedder{ |
| 40 | client: openai.NewClient(opts...), | 39 | client: openai.NewClient(opts...), |
| 41 | model: cfg.Model, | 40 | model: cfg.Model, |
| 42 | dimension: cfg.Dimension, | 41 | dimension: cfg.Dimension, |
| 43 | } | 42 | } |
| 44 | } | 43 | } |
| 45 | 44 | ||
| 46 | -func (e *EmbedderClient) Dimension() int { | 45 | +func (e *embedder) Dimension() int { |
| 47 | return e.dimension | 46 | return e.dimension |
| 48 | } | 47 | } |
| 49 | 48 | ||
| 50 | -func (e *EmbedderClient) Embed(ctx context.Context, texts []string, instruction string) ([][]float32, error) { | 49 | +func (e *embedder) Embed(ctx context.Context, texts []string, instruction string) ([][]float32, error) { |
| 51 | inputs := texts | 50 | inputs := texts |
| 52 | if instruction != "" { | 51 | if instruction != "" { |
| 53 | inputs = make([]string, len(texts)) | 52 | inputs = make([]string, len(texts)) |
| @@ -75,11 +74,11 @@ func (e *EmbedderClient) Embed(ctx context.Context, texts []string, instruction | |||
| 75 | return embeddings, nil | 74 | return embeddings, nil |
| 76 | } | 75 | } |
| 77 | 76 | ||
| 78 | -func avgEmbeddings(blobs [][]byte, dim int) ([]byte, error) { | 77 | +func AvgEmbeddings(blobs [][]byte, dim int) ([]byte, error) { |
| 79 | sum := make([]float32, dim) | 78 | sum := make([]float32, dim) |
| 80 | count := 0 | 79 | count := 0 |
| 81 | for _, blob := range blobs { | 80 | for _, blob := range blobs { |
| 82 | - v := bytesToFloat32s(blob, dim) | 81 | + v := BytesToFloat32s(blob, dim) |
| 83 | if v == nil { | 82 | if v == nil { |
| 84 | continue | 83 | continue |
| 85 | } | 84 | } |
| @@ -97,7 +96,7 @@ func avgEmbeddings(blobs [][]byte, dim int) ([]byte, error) { | |||
| 97 | return vec.SerializeFloat32(sum) | 96 | return vec.SerializeFloat32(sum) |
| 98 | } | 97 | } |
| 99 | 98 | ||
| 100 | -func bytesToFloat32s(data []byte, expectedDim int) []float32 { | 99 | +func BytesToFloat32s(data []byte, expectedDim int) []float32 { |
| 101 | if len(data) != expectedDim*4 { | 100 | if len(data) != expectedDim*4 { |
| 102 | return nil | 101 | return nil |
| 103 | } | 102 | } |
renamed
internal/ai/llm.go +10 -6 | similarity index 84% | ||
| rename from internal/cluster/llm.go | ||
| rename to internal/ai/llm.go | ||
| @@ -1,4 +1,4 @@ | ||
| 1 | -package cluster | |
| 1 | +package ai | |
| 2 | 2 | |
| 3 | 3 | import ( |
| 4 | 4 | "context" |
| @@ -11,18 +11,22 @@ import ( | ||
| 11 | 11 | "pkg.rbrt.fr/glean/internal/langdetect" |
| 12 | 12 | ) |
| 13 | 13 | |
| 14 | -type LLMClient struct { | |
| 14 | +type TextModel interface { | |
| 15 | + DetectLanguages(ctx context.Context, texts []string) ([]string, error) | |
| 16 | +} | |
| 17 | + | |
| 18 | +type llm struct { | |
| 15 | 19 | client openai.Client |
| 16 | 20 | model string |
| 17 | 21 | } |
| 18 | 22 | |
| 19 | -type LLMClientConfig struct { | |
| 23 | +type LLMConfig struct { | |
| 20 | 24 | BaseURL string |
| 21 | 25 | APIKey string |
| 22 | 26 | Model string |
| 23 | 27 | } |
| 24 | 28 | |
| 25 | -func NewLLMClient(cfg LLMClientConfig) *LLMClient { | |
| 29 | +func NewLLM(cfg LLMConfig) TextModel { | |
| 26 | 30 | opts := []option.RequestOption{} |
| 27 | 31 | if cfg.BaseURL != "" { |
| 28 | 32 | opts = append(opts, option.WithBaseURL(cfg.BaseURL)) |
| @@ -30,13 +34,13 @@ func NewLLMClient(cfg LLMClientConfig) *LLMClient { | ||
| 30 | 34 | if cfg.APIKey != "" { |
| 31 | 35 | opts = append(opts, option.WithAPIKey(cfg.APIKey)) |
| 32 | 36 | } |
| 33 | - return &LLMClient{ | |
| 37 | + return &llm{ | |
| 34 | 38 | client: openai.NewClient(opts...), |
| 35 | 39 | model: cfg.Model, |
| 36 | 40 | } |
| 37 | 41 | } |
| 38 | 42 | |
| 39 | -func (c *LLMClient) DetectLanguages(ctx context.Context, texts []string) ([]string, error) { | |
| 43 | +func (c *llm) DetectLanguages(ctx context.Context, texts []string) ([]string, error) { | |
| 40 | 44 | var b strings.Builder |
| 41 | 45 | b.WriteString("For each text below, respond with ONLY the ISO 639-1 language code (e.g. en, fr, de, es, pt, it, ru, ja, zh, ko, ar). One code per line, same order as input. If uncertain, respond with 'en'.\n\n") |
| 42 | 46 | for i, t := range texts { |
| similarity index 84% | |||
| rename from internal/cluster/llm.go | |||
| rename to internal/ai/llm.go | |||
| @@ -1,4 +1,4 @@ | |||
| 1 | -package cluster | 1 | +package ai |
| 2 | 2 | ||
| 3 | import ( | 3 | import ( |
| 4 | "context" | 4 | "context" |
| @@ -11,18 +11,22 @@ import ( | |||
| 11 | "pkg.rbrt.fr/glean/internal/langdetect" | 11 | "pkg.rbrt.fr/glean/internal/langdetect" |
| 12 | ) | 12 | ) |
| 13 | 13 | ||
| 14 | -type LLMClient struct { | 14 | +type TextModel interface { |
| 15 | + DetectLanguages(ctx context.Context, texts []string) ([]string, error) | ||
| 16 | +} | ||
| 17 | + | ||
| 18 | +type llm struct { | ||
| 15 | client openai.Client | 19 | client openai.Client |
| 16 | model string | 20 | model string |
| 17 | } | 21 | } |
| 18 | 22 | ||
| 19 | -type LLMClientConfig struct { | 23 | +type LLMConfig struct { |
| 20 | BaseURL string | 24 | BaseURL string |
| 21 | APIKey string | 25 | APIKey string |
| 22 | Model string | 26 | Model string |
| 23 | } | 27 | } |
| 24 | 28 | ||
| 25 | -func NewLLMClient(cfg LLMClientConfig) *LLMClient { | 29 | +func NewLLM(cfg LLMConfig) TextModel { |
| 26 | opts := []option.RequestOption{} | 30 | opts := []option.RequestOption{} |
| 27 | if cfg.BaseURL != "" { | 31 | if cfg.BaseURL != "" { |
| 28 | opts = append(opts, option.WithBaseURL(cfg.BaseURL)) | 32 | opts = append(opts, option.WithBaseURL(cfg.BaseURL)) |
| @@ -30,13 +34,13 @@ func NewLLMClient(cfg LLMClientConfig) *LLMClient { | |||
| 30 | if cfg.APIKey != "" { | 34 | if cfg.APIKey != "" { |
| 31 | opts = append(opts, option.WithAPIKey(cfg.APIKey)) | 35 | opts = append(opts, option.WithAPIKey(cfg.APIKey)) |
| 32 | } | 36 | } |
| 33 | - return &LLMClient{ | 37 | + return &llm{ |
| 34 | client: openai.NewClient(opts...), | 38 | client: openai.NewClient(opts...), |
| 35 | model: cfg.Model, | 39 | model: cfg.Model, |
| 36 | } | 40 | } |
| 37 | } | 41 | } |
| 38 | 42 | ||
| 39 | -func (c *LLMClient) DetectLanguages(ctx context.Context, texts []string) ([]string, error) { | 43 | +func (c *llm) DetectLanguages(ctx context.Context, texts []string) ([]string, error) { |
| 40 | var b strings.Builder | 44 | var b strings.Builder |
| 41 | b.WriteString("For each text below, respond with ONLY the ISO 639-1 language code (e.g. en, fr, de, es, pt, it, ru, ja, zh, ko, ar). One code per line, same order as input. If uncertain, respond with 'en'.\n\n") | 45 | b.WriteString("For each text below, respond with ONLY the ISO 639-1 language code (e.g. en, fr, de, es, pt, it, ru, ja, zh, ko, ar). One code per line, same order as input. If uncertain, respond with 'en'.\n\n") |
| 42 | for i, t := range texts { | 46 | for i, t := range texts { |
modified
internal/cluster/article.go +3 -1 | @@ -8,6 +8,8 @@ import ( | ||
| 8 | 8 | "strings" |
| 9 | 9 | |
| 10 | 10 | vec "github.com/asg017/sqlite-vec-go-bindings/cgo" |
| 11 | + | |
| 12 | + "pkg.rbrt.fr/glean/internal/ai" | |
| 11 | 13 | ) |
| 12 | 14 | |
| 13 | 15 | // embedBatchSize caps how many texts are sent in a single embedding API call. |
| @@ -194,7 +196,7 @@ func (e *Engine) populateContentBoost(ctx context.Context, conn *sql.Conn, userD | ||
| 194 | 196 | return nil |
| 195 | 197 | } |
| 196 | 198 | |
| 197 | - queryBlob, err := avgEmbeddings(blobs, dim) | |
| 199 | + queryBlob, err := ai.AvgEmbeddings(blobs, dim) | |
| 198 | 200 | if err != nil { |
| 199 | 201 | return fmt.Errorf("serialize query vector: %w", err) |
| 200 | 202 | } |
| @@ -8,6 +8,8 @@ import ( | |||
| 8 | "strings" | 8 | "strings" |
| 9 | 9 | ||
| 10 | vec "github.com/asg017/sqlite-vec-go-bindings/cgo" | 10 | vec "github.com/asg017/sqlite-vec-go-bindings/cgo" |
| 11 | + | ||
| 12 | + "pkg.rbrt.fr/glean/internal/ai" | ||
| 11 | ) | 13 | ) |
| 12 | 14 | ||
| 13 | // embedBatchSize caps how many texts are sent in a single embedding API call. | 15 | // embedBatchSize caps how many texts are sent in a single embedding API call. |
| @@ -194,7 +196,7 @@ func (e *Engine) populateContentBoost(ctx context.Context, conn *sql.Conn, userD | |||
| 194 | return nil | 196 | return nil |
| 195 | } | 197 | } |
| 196 | 198 | ||
| 197 | - queryBlob, err := avgEmbeddings(blobs, dim) | 199 | + queryBlob, err := ai.AvgEmbeddings(blobs, dim) |
| 198 | if err != nil { | 200 | if err != nil { |
| 199 | return fmt.Errorf("serialize query vector: %w", err) | 201 | return fmt.Errorf("serialize query vector: %w", err) |
| 200 | } | 202 | } |
modified
internal/cluster/cron.go +1 -1 | @@ -59,7 +59,7 @@ func (c *Cron) Run(ctx context.Context) error { | ||
| 59 | 59 | if err := c.engine.ComputeFollowDistances(ctx); err != nil { |
| 60 | 60 | c.logger.Error("follow distances failed", "error", err) |
| 61 | 61 | } |
| 62 | - if err := c.engine.AutoDismissStale(ctx, 5, 5); err != nil { | |
| 62 | + if err := c.engine.feedback.AutoDismissStale(ctx, 5, 5); err != nil { | |
| 63 | 63 | c.engine.logger.Error("auto dismiss failed", "error", err) |
| 64 | 64 | } |
| 65 | 65 | c.engine.mu.Unlock() |
| @@ -59,7 +59,7 @@ func (c *Cron) Run(ctx context.Context) error { | |||
| 59 | if err := c.engine.ComputeFollowDistances(ctx); err != nil { | 59 | if err := c.engine.ComputeFollowDistances(ctx); err != nil { |
| 60 | c.logger.Error("follow distances failed", "error", err) | 60 | c.logger.Error("follow distances failed", "error", err) |
| 61 | } | 61 | } |
| 62 | - if err := c.engine.AutoDismissStale(ctx, 5, 5); err != nil { | 62 | + if err := c.engine.feedback.AutoDismissStale(ctx, 5, 5); err != nil { |
| 63 | c.engine.logger.Error("auto dismiss failed", "error", err) | 63 | c.engine.logger.Error("auto dismiss failed", "error", err) |
| 64 | } | 64 | } |
| 65 | c.engine.mu.Unlock() | 65 | c.engine.mu.Unlock() |
modified
internal/cluster/jaccard.go +7 -3 | @@ -10,7 +10,9 @@ import ( | ||
| 10 | 10 | |
| 11 | 11 | "github.com/hashicorp/golang-lru/v2/expirable" |
| 12 | 12 | |
| 13 | + "pkg.rbrt.fr/glean/internal/ai" | |
| 13 | 14 | "pkg.rbrt.fr/glean/internal/db" |
| 15 | + "pkg.rbrt.fr/glean/internal/feedback" | |
| 14 | 16 | ) |
| 15 | 17 | |
| 16 | 18 | // Config controls weights used during similarity computation (feed similarity |
| @@ -44,8 +46,9 @@ type Engine struct { | ||
| 44 | 46 | mu sync.Mutex |
| 45 | 47 | config Config |
| 46 | 48 | |
| 47 | - embedder Embedder | |
| 48 | - llm *LLMClient | |
| 49 | + embedder ai.Embedder | |
| 50 | + llm ai.TextModel | |
| 51 | + feedback *feedback.Service | |
| 49 | 52 | |
| 50 | 53 | feedCache *expirable.LRU[string, []*FeedRecommendation] |
| 51 | 54 | peopleCache *expirable.LRU[string, []*PersonRecommendation] |
| @@ -57,7 +60,7 @@ type Engine struct { | ||
| 57 | 60 | // recCacheSize is the maximum number of recommendations to cache per user. |
| 58 | 61 | const recCacheSize = 512 |
| 59 | 62 | |
| 60 | -func NewEngine(sqlDB *sql.DB, articles *db.ArticleStore, embedder Embedder, llm *LLMClient, logger *slog.Logger, cacheTTL time.Duration, config Config) *Engine { | |
| 63 | +func NewEngine(sqlDB *sql.DB, articles *db.ArticleStore, embedder ai.Embedder, llm ai.TextModel, fb *feedback.Service, logger *slog.Logger, cacheTTL time.Duration, config Config) *Engine { | |
| 61 | 64 | return &Engine{ |
| 62 | 65 | db: sqlDB, |
| 63 | 66 | articles: articles, |
| @@ -65,6 +68,7 @@ func NewEngine(sqlDB *sql.DB, articles *db.ArticleStore, embedder Embedder, llm | ||
| 65 | 68 | config: config, |
| 66 | 69 | embedder: embedder, |
| 67 | 70 | llm: llm, |
| 71 | + feedback: fb, | |
| 68 | 72 | feedCache: expirable.NewLRU[string, []*FeedRecommendation](recCacheSize, nil, cacheTTL), |
| 69 | 73 | peopleCache: expirable.NewLRU[string, []*PersonRecommendation](recCacheSize, nil, cacheTTL), |
| 70 | 74 | articleCache: expirable.NewLRU[string, []*ArticleRecommendation](recCacheSize, nil, cacheTTL), |
| @@ -10,7 +10,9 @@ import ( | |||
| 10 | 10 | ||
| 11 | "github.com/hashicorp/golang-lru/v2/expirable" | 11 | "github.com/hashicorp/golang-lru/v2/expirable" |
| 12 | 12 | ||
| 13 | + "pkg.rbrt.fr/glean/internal/ai" | ||
| 13 | "pkg.rbrt.fr/glean/internal/db" | 14 | "pkg.rbrt.fr/glean/internal/db" |
| 15 | + "pkg.rbrt.fr/glean/internal/feedback" | ||
| 14 | ) | 16 | ) |
| 15 | 17 | ||
| 16 | // Config controls weights used during similarity computation (feed similarity | 18 | // Config controls weights used during similarity computation (feed similarity |
| @@ -44,8 +46,9 @@ type Engine struct { | |||
| 44 | mu sync.Mutex | 46 | mu sync.Mutex |
| 45 | config Config | 47 | config Config |
| 46 | 48 | ||
| 47 | - embedder Embedder | 49 | + embedder ai.Embedder |
| 48 | - llm *LLMClient | 50 | + llm ai.TextModel |
| 51 | + feedback *feedback.Service | ||
| 49 | 52 | ||
| 50 | feedCache *expirable.LRU[string, []*FeedRecommendation] | 53 | feedCache *expirable.LRU[string, []*FeedRecommendation] |
| 51 | peopleCache *expirable.LRU[string, []*PersonRecommendation] | 54 | peopleCache *expirable.LRU[string, []*PersonRecommendation] |
| @@ -57,7 +60,7 @@ type Engine struct { | |||
| 57 | // recCacheSize is the maximum number of recommendations to cache per user. | 60 | // recCacheSize is the maximum number of recommendations to cache per user. |
| 58 | const recCacheSize = 512 | 61 | const recCacheSize = 512 |
| 59 | 62 | ||
| 60 | -func NewEngine(sqlDB *sql.DB, articles *db.ArticleStore, embedder Embedder, llm *LLMClient, logger *slog.Logger, cacheTTL time.Duration, config Config) *Engine { | 63 | +func NewEngine(sqlDB *sql.DB, articles *db.ArticleStore, embedder ai.Embedder, llm ai.TextModel, fb *feedback.Service, logger *slog.Logger, cacheTTL time.Duration, config Config) *Engine { |
| 61 | return &Engine{ | 64 | return &Engine{ |
| 62 | db: sqlDB, | 65 | db: sqlDB, |
| 63 | articles: articles, | 66 | articles: articles, |
| @@ -65,6 +68,7 @@ func NewEngine(sqlDB *sql.DB, articles *db.ArticleStore, embedder Embedder, llm | |||
| 65 | config: config, | 68 | config: config, |
| 66 | embedder: embedder, | 69 | embedder: embedder, |
| 67 | llm: llm, | 70 | llm: llm, |
| 71 | + feedback: fb, | ||
| 68 | feedCache: expirable.NewLRU[string, []*FeedRecommendation](recCacheSize, nil, cacheTTL), | 72 | feedCache: expirable.NewLRU[string, []*FeedRecommendation](recCacheSize, nil, cacheTTL), |
| 69 | peopleCache: expirable.NewLRU[string, []*PersonRecommendation](recCacheSize, nil, cacheTTL), | 73 | peopleCache: expirable.NewLRU[string, []*PersonRecommendation](recCacheSize, nil, cacheTTL), |
| 70 | articleCache: expirable.NewLRU[string, []*ArticleRecommendation](recCacheSize, nil, cacheTTL), | 74 | articleCache: expirable.NewLRU[string, []*ArticleRecommendation](recCacheSize, nil, cacheTTL), |
modified
internal/cluster/jaccard_test.go +21 -20 | @@ -12,6 +12,7 @@ import ( | ||
| 12 | 12 | |
| 13 | 13 | vec "github.com/asg017/sqlite-vec-go-bindings/cgo" |
| 14 | 14 | "pkg.rbrt.fr/glean/internal/db" |
| 15 | + "pkg.rbrt.fr/glean/internal/feedback" | |
| 15 | 16 | |
| 16 | 17 | "gotest.tools/v3/assert" |
| 17 | 18 | ) |
| @@ -95,7 +96,7 @@ func seedFollowData(t *testing.T, ctx context.Context, dbs *db.Store) { | ||
| 95 | 96 | } |
| 96 | 97 | |
| 97 | 98 | func newTestEngine(dbs *db.Store) *Engine { |
| 98 | - return NewEngine(dbs.SQLDB(), dbs.Articles, NewMockEmbedder(8), nil, slog.Default(), time.Hour, DefaultConfig()) | |
| 99 | + return NewEngine(dbs.SQLDB(), dbs.Articles, NewMockEmbedder(8), nil, feedback.NewService(dbs.SQLDB()), slog.Default(), time.Hour, DefaultConfig()) | |
| 99 | 100 | } |
| 100 | 101 | |
| 101 | 102 | func TestComputeFeedSimilarity(t *testing.T) { |
| @@ -182,7 +183,7 @@ func TestDismissedFeedsExcluded(t *testing.T) { | ||
| 182 | 183 | assert.NilError(t, engine.ComputeFeedSimilarity(ctx)) |
| 183 | 184 | assert.NilError(t, engine.ComputeUserSimilarity(ctx)) |
| 184 | 185 | |
| 185 | - assert.NilError(t, engine.DismissFeed(ctx, "did:test:carol", "https://a.com/feed", "not_interested")) | |
| 186 | + assert.NilError(t, engine.feedback.DismissFeed(ctx, "did:test:carol", "https://a.com/feed", "not_interested")) | |
| 186 | 187 | |
| 187 | 188 | recs, err := engine.GetFeedRecommendations(ctx, "did:test:carol", 10) |
| 188 | 189 | assert.NilError(t, err) |
| @@ -200,13 +201,13 @@ func TestIsFeedDismissed(t *testing.T) { | ||
| 200 | 201 | |
| 201 | 202 | engine := newTestEngine(dbs) |
| 202 | 203 | |
| 203 | - dismissed, err := engine.IsFeedDismissed(ctx, "did:test:alice", "https://a.com/feed") | |
| 204 | + dismissed, err := engine.feedback.IsFeedDismissed(ctx, "did:test:alice", "https://a.com/feed") | |
| 204 | 205 | assert.NilError(t, err) |
| 205 | 206 | assert.Assert(t, !dismissed, "feed should not be dismissed initially") |
| 206 | 207 | |
| 207 | - assert.NilError(t, engine.DismissFeed(ctx, "did:test:alice", "https://a.com/feed", "not_interested")) | |
| 208 | + assert.NilError(t, engine.feedback.DismissFeed(ctx, "did:test:alice", "https://a.com/feed", "not_interested")) | |
| 208 | 209 | |
| 209 | - dismissed, err = engine.IsFeedDismissed(ctx, "did:test:alice", "https://a.com/feed") | |
| 210 | + dismissed, err = engine.feedback.IsFeedDismissed(ctx, "did:test:alice", "https://a.com/feed") | |
| 210 | 211 | assert.NilError(t, err) |
| 211 | 212 | assert.Assert(t, dismissed, "feed should be dismissed after dismiss call") |
| 212 | 213 | } |
| @@ -218,18 +219,18 @@ func TestRecordImpressions(t *testing.T) { | ||
| 218 | 219 | |
| 219 | 220 | engine := newTestEngine(dbs) |
| 220 | 221 | |
| 221 | - impressions := []Impression{ | |
| 222 | + impressions := []feedback.Impression{ | |
| 222 | 223 | {TargetType: "feed", TargetID: "https://a.com/feed"}, |
| 223 | 224 | {TargetType: "feed", TargetID: "https://b.com/feed"}, |
| 224 | 225 | } |
| 225 | - assert.NilError(t, engine.RecordImpressions(ctx, "did:test:alice", impressions)) | |
| 226 | + assert.NilError(t, engine.feedback.RecordImpressions(ctx, "did:test:alice", impressions)) | |
| 226 | 227 | |
| 227 | 228 | var count int |
| 228 | 229 | assert.NilError(t, dbs.SQLDB().QueryRowContext(ctx, |
| 229 | 230 | `SELECT COUNT(*) FROM main.recommendation_impressions WHERE user_did = 'did:test:alice'`).Scan(&count)) |
| 230 | 231 | assert.Equal(t, count, 2) |
| 231 | 232 | |
| 232 | - assert.NilError(t, engine.RecordImpressions(ctx, "did:test:alice", impressions)) | |
| 233 | + assert.NilError(t, engine.feedback.RecordImpressions(ctx, "did:test:alice", impressions)) | |
| 233 | 234 | |
| 234 | 235 | var shownCount int |
| 235 | 236 | assert.NilError(t, dbs.SQLDB().QueryRowContext(ctx, |
| @@ -244,10 +245,10 @@ func TestMarkImpressionActed(t *testing.T) { | ||
| 244 | 245 | |
| 245 | 246 | engine := newTestEngine(dbs) |
| 246 | 247 | |
| 247 | - impressions := []Impression{{TargetType: "feed", TargetID: "https://a.com/feed"}} | |
| 248 | - assert.NilError(t, engine.RecordImpressions(ctx, "did:test:alice", impressions)) | |
| 248 | + impressions := []feedback.Impression{{TargetType: "feed", TargetID: "https://a.com/feed"}} | |
| 249 | + assert.NilError(t, engine.feedback.RecordImpressions(ctx, "did:test:alice", impressions)) | |
| 249 | 250 | |
| 250 | - assert.NilError(t, engine.MarkImpressionActed(ctx, "did:test:alice", "feed", "https://a.com/feed")) | |
| 251 | + assert.NilError(t, engine.feedback.MarkImpressionActed(ctx, "did:test:alice", "feed", "https://a.com/feed")) | |
| 251 | 252 | |
| 252 | 253 | var acted bool |
| 253 | 254 | assert.NilError(t, dbs.SQLDB().QueryRowContext(ctx, |
| @@ -313,9 +314,9 @@ func TestAutoDismissStale(t *testing.T) { | ||
| 313 | 314 | `) |
| 314 | 315 | assert.NilError(t, err) |
| 315 | 316 | |
| 316 | - assert.NilError(t, engine.AutoDismissStale(ctx, 5, 5)) | |
| 317 | + assert.NilError(t, engine.feedback.AutoDismissStale(ctx, 5, 5)) | |
| 317 | 318 | |
| 318 | - dismissed, err := engine.IsFeedDismissed(ctx, "did:test:alice", "https://stale.com/feed") | |
| 319 | + dismissed, err := engine.feedback.IsFeedDismissed(ctx, "did:test:alice", "https://stale.com/feed") | |
| 319 | 320 | assert.NilError(t, err) |
| 320 | 321 | assert.Assert(t, dismissed, "stale recommendation should be auto-dismissed") |
| 321 | 322 | } |
| @@ -333,9 +334,9 @@ func TestAutoDismissStale_DoesNotDismissRecent(t *testing.T) { | ||
| 333 | 334 | `) |
| 334 | 335 | assert.NilError(t, err) |
| 335 | 336 | |
| 336 | - assert.NilError(t, engine.AutoDismissStale(ctx, 5, 5)) | |
| 337 | + assert.NilError(t, engine.feedback.AutoDismissStale(ctx, 5, 5)) | |
| 337 | 338 | |
| 338 | - dismissed, err := engine.IsFeedDismissed(ctx, "did:test:alice", "https://recent.com/feed") | |
| 339 | + dismissed, err := engine.feedback.IsFeedDismissed(ctx, "did:test:alice", "https://recent.com/feed") | |
| 339 | 340 | assert.NilError(t, err) |
| 340 | 341 | assert.Assert(t, !dismissed, "recent impression should not be auto-dismissed") |
| 341 | 342 | } |
| @@ -353,9 +354,9 @@ func TestAutoDismissStale_DoesNotDismissActed(t *testing.T) { | ||
| 353 | 354 | `) |
| 354 | 355 | assert.NilError(t, err) |
| 355 | 356 | |
| 356 | - assert.NilError(t, engine.AutoDismissStale(ctx, 5, 5)) | |
| 357 | + assert.NilError(t, engine.feedback.AutoDismissStale(ctx, 5, 5)) | |
| 357 | 358 | |
| 358 | - dismissed, err := engine.IsFeedDismissed(ctx, "did:test:alice", "https://acted.com/feed") | |
| 359 | + dismissed, err := engine.feedback.IsFeedDismissed(ctx, "did:test:alice", "https://acted.com/feed") | |
| 359 | 360 | assert.NilError(t, err) |
| 360 | 361 | assert.Assert(t, !dismissed, "acted recommendation should not be auto-dismissed") |
| 361 | 362 | } |
| @@ -503,7 +504,7 @@ func TestDismissArticle(t *testing.T) { | ||
| 503 | 504 | |
| 504 | 505 | engine := newTestEngine(dbs) |
| 505 | 506 | |
| 506 | - assert.NilError(t, engine.DismissArticle(ctx, "did:test:alice", "https://a.com/article1", "not_interested")) | |
| 507 | + assert.NilError(t, engine.feedback.DismissArticle(ctx, "did:test:alice", "https://a.com/article1", "not_interested")) | |
| 507 | 508 | |
| 508 | 509 | var count int |
| 509 | 510 | assert.NilError(t, dbs.SQLDB().QueryRowContext(ctx, |
| @@ -531,8 +532,8 @@ func TestDismissFeed_Idempotent(t *testing.T) { | ||
| 531 | 532 | |
| 532 | 533 | engine := newTestEngine(dbs) |
| 533 | 534 | |
| 534 | - assert.NilError(t, engine.DismissFeed(ctx, "did:test:alice", "https://a.com/feed", "reason1")) | |
| 535 | - assert.NilError(t, engine.DismissFeed(ctx, "did:test:alice", "https://a.com/feed", "reason2")) | |
| 535 | + assert.NilError(t, engine.feedback.DismissFeed(ctx, "did:test:alice", "https://a.com/feed", "reason1")) | |
| 536 | + assert.NilError(t, engine.feedback.DismissFeed(ctx, "did:test:alice", "https://a.com/feed", "reason2")) | |
| 536 | 537 | |
| 537 | 538 | var count int |
| 538 | 539 | assert.NilError(t, dbs.SQLDB().QueryRowContext(ctx, |
| @@ -12,6 +12,7 @@ import ( | |||
| 12 | 12 | ||
| 13 | vec "github.com/asg017/sqlite-vec-go-bindings/cgo" | 13 | vec "github.com/asg017/sqlite-vec-go-bindings/cgo" |
| 14 | "pkg.rbrt.fr/glean/internal/db" | 14 | "pkg.rbrt.fr/glean/internal/db" |
| 15 | + "pkg.rbrt.fr/glean/internal/feedback" | ||
| 15 | 16 | ||
| 16 | "gotest.tools/v3/assert" | 17 | "gotest.tools/v3/assert" |
| 17 | ) | 18 | ) |
| @@ -95,7 +96,7 @@ func seedFollowData(t *testing.T, ctx context.Context, dbs *db.Store) { | |||
| 95 | } | 96 | } |
| 96 | 97 | ||
| 97 | func newTestEngine(dbs *db.Store) *Engine { | 98 | func newTestEngine(dbs *db.Store) *Engine { |
| 98 | - return NewEngine(dbs.SQLDB(), dbs.Articles, NewMockEmbedder(8), nil, slog.Default(), time.Hour, DefaultConfig()) | 99 | + return NewEngine(dbs.SQLDB(), dbs.Articles, NewMockEmbedder(8), nil, feedback.NewService(dbs.SQLDB()), slog.Default(), time.Hour, DefaultConfig()) |
| 99 | } | 100 | } |
| 100 | 101 | ||
| 101 | func TestComputeFeedSimilarity(t *testing.T) { | 102 | func TestComputeFeedSimilarity(t *testing.T) { |
| @@ -182,7 +183,7 @@ func TestDismissedFeedsExcluded(t *testing.T) { | |||
| 182 | assert.NilError(t, engine.ComputeFeedSimilarity(ctx)) | 183 | assert.NilError(t, engine.ComputeFeedSimilarity(ctx)) |
| 183 | assert.NilError(t, engine.ComputeUserSimilarity(ctx)) | 184 | assert.NilError(t, engine.ComputeUserSimilarity(ctx)) |
| 184 | 185 | ||
| 185 | - assert.NilError(t, engine.DismissFeed(ctx, "did:test:carol", "https://a.com/feed", "not_interested")) | 186 | + assert.NilError(t, engine.feedback.DismissFeed(ctx, "did:test:carol", "https://a.com/feed", "not_interested")) |
| 186 | 187 | ||
| 187 | recs, err := engine.GetFeedRecommendations(ctx, "did:test:carol", 10) | 188 | recs, err := engine.GetFeedRecommendations(ctx, "did:test:carol", 10) |
| 188 | assert.NilError(t, err) | 189 | assert.NilError(t, err) |
| @@ -200,13 +201,13 @@ func TestIsFeedDismissed(t *testing.T) { | |||
| 200 | 201 | ||
| 201 | engine := newTestEngine(dbs) | 202 | engine := newTestEngine(dbs) |
| 202 | 203 | ||
| 203 | - dismissed, err := engine.IsFeedDismissed(ctx, "did:test:alice", "https://a.com/feed") | 204 | + dismissed, err := engine.feedback.IsFeedDismissed(ctx, "did:test:alice", "https://a.com/feed") |
| 204 | assert.NilError(t, err) | 205 | assert.NilError(t, err) |
| 205 | assert.Assert(t, !dismissed, "feed should not be dismissed initially") | 206 | assert.Assert(t, !dismissed, "feed should not be dismissed initially") |
| 206 | 207 | ||
| 207 | - assert.NilError(t, engine.DismissFeed(ctx, "did:test:alice", "https://a.com/feed", "not_interested")) | 208 | + assert.NilError(t, engine.feedback.DismissFeed(ctx, "did:test:alice", "https://a.com/feed", "not_interested")) |
| 208 | 209 | ||
| 209 | - dismissed, err = engine.IsFeedDismissed(ctx, "did:test:alice", "https://a.com/feed") | 210 | + dismissed, err = engine.feedback.IsFeedDismissed(ctx, "did:test:alice", "https://a.com/feed") |
| 210 | assert.NilError(t, err) | 211 | assert.NilError(t, err) |
| 211 | assert.Assert(t, dismissed, "feed should be dismissed after dismiss call") | 212 | assert.Assert(t, dismissed, "feed should be dismissed after dismiss call") |
| 212 | } | 213 | } |
| @@ -218,18 +219,18 @@ func TestRecordImpressions(t *testing.T) { | |||
| 218 | 219 | ||
| 219 | engine := newTestEngine(dbs) | 220 | engine := newTestEngine(dbs) |
| 220 | 221 | ||
| 221 | - impressions := []Impression{ | 222 | + impressions := []feedback.Impression{ |
| 222 | {TargetType: "feed", TargetID: "https://a.com/feed"}, | 223 | {TargetType: "feed", TargetID: "https://a.com/feed"}, |
| 223 | {TargetType: "feed", TargetID: "https://b.com/feed"}, | 224 | {TargetType: "feed", TargetID: "https://b.com/feed"}, |
| 224 | } | 225 | } |
| 225 | - assert.NilError(t, engine.RecordImpressions(ctx, "did:test:alice", impressions)) | 226 | + assert.NilError(t, engine.feedback.RecordImpressions(ctx, "did:test:alice", impressions)) |
| 226 | 227 | ||
| 227 | var count int | 228 | var count int |
| 228 | assert.NilError(t, dbs.SQLDB().QueryRowContext(ctx, | 229 | assert.NilError(t, dbs.SQLDB().QueryRowContext(ctx, |
| 229 | `SELECT COUNT(*) FROM main.recommendation_impressions WHERE user_did = 'did:test:alice'`).Scan(&count)) | 230 | `SELECT COUNT(*) FROM main.recommendation_impressions WHERE user_did = 'did:test:alice'`).Scan(&count)) |
| 230 | assert.Equal(t, count, 2) | 231 | assert.Equal(t, count, 2) |
| 231 | 232 | ||
| 232 | - assert.NilError(t, engine.RecordImpressions(ctx, "did:test:alice", impressions)) | 233 | + assert.NilError(t, engine.feedback.RecordImpressions(ctx, "did:test:alice", impressions)) |
| 233 | 234 | ||
| 234 | var shownCount int | 235 | var shownCount int |
| 235 | assert.NilError(t, dbs.SQLDB().QueryRowContext(ctx, | 236 | assert.NilError(t, dbs.SQLDB().QueryRowContext(ctx, |
| @@ -244,10 +245,10 @@ func TestMarkImpressionActed(t *testing.T) { | |||
| 244 | 245 | ||
| 245 | engine := newTestEngine(dbs) | 246 | engine := newTestEngine(dbs) |
| 246 | 247 | ||
| 247 | - impressions := []Impression{{TargetType: "feed", TargetID: "https://a.com/feed"}} | 248 | + impressions := []feedback.Impression{{TargetType: "feed", TargetID: "https://a.com/feed"}} |
| 248 | - assert.NilError(t, engine.RecordImpressions(ctx, "did:test:alice", impressions)) | 249 | + assert.NilError(t, engine.feedback.RecordImpressions(ctx, "did:test:alice", impressions)) |
| 249 | 250 | ||
| 250 | - assert.NilError(t, engine.MarkImpressionActed(ctx, "did:test:alice", "feed", "https://a.com/feed")) | 251 | + assert.NilError(t, engine.feedback.MarkImpressionActed(ctx, "did:test:alice", "feed", "https://a.com/feed")) |
| 251 | 252 | ||
| 252 | var acted bool | 253 | var acted bool |
| 253 | assert.NilError(t, dbs.SQLDB().QueryRowContext(ctx, | 254 | assert.NilError(t, dbs.SQLDB().QueryRowContext(ctx, |
| @@ -313,9 +314,9 @@ func TestAutoDismissStale(t *testing.T) { | |||
| 313 | `) | 314 | `) |
| 314 | assert.NilError(t, err) | 315 | assert.NilError(t, err) |
| 315 | 316 | ||
| 316 | - assert.NilError(t, engine.AutoDismissStale(ctx, 5, 5)) | 317 | + assert.NilError(t, engine.feedback.AutoDismissStale(ctx, 5, 5)) |
| 317 | 318 | ||
| 318 | - dismissed, err := engine.IsFeedDismissed(ctx, "did:test:alice", "https://stale.com/feed") | 319 | + dismissed, err := engine.feedback.IsFeedDismissed(ctx, "did:test:alice", "https://stale.com/feed") |
| 319 | assert.NilError(t, err) | 320 | assert.NilError(t, err) |
| 320 | assert.Assert(t, dismissed, "stale recommendation should be auto-dismissed") | 321 | assert.Assert(t, dismissed, "stale recommendation should be auto-dismissed") |
| 321 | } | 322 | } |
| @@ -333,9 +334,9 @@ func TestAutoDismissStale_DoesNotDismissRecent(t *testing.T) { | |||
| 333 | `) | 334 | `) |
| 334 | assert.NilError(t, err) | 335 | assert.NilError(t, err) |
| 335 | 336 | ||
| 336 | - assert.NilError(t, engine.AutoDismissStale(ctx, 5, 5)) | 337 | + assert.NilError(t, engine.feedback.AutoDismissStale(ctx, 5, 5)) |
| 337 | 338 | ||
| 338 | - dismissed, err := engine.IsFeedDismissed(ctx, "did:test:alice", "https://recent.com/feed") | 339 | + dismissed, err := engine.feedback.IsFeedDismissed(ctx, "did:test:alice", "https://recent.com/feed") |
| 339 | assert.NilError(t, err) | 340 | assert.NilError(t, err) |
| 340 | assert.Assert(t, !dismissed, "recent impression should not be auto-dismissed") | 341 | assert.Assert(t, !dismissed, "recent impression should not be auto-dismissed") |
| 341 | } | 342 | } |
| @@ -353,9 +354,9 @@ func TestAutoDismissStale_DoesNotDismissActed(t *testing.T) { | |||
| 353 | `) | 354 | `) |
| 354 | assert.NilError(t, err) | 355 | assert.NilError(t, err) |
| 355 | 356 | ||
| 356 | - assert.NilError(t, engine.AutoDismissStale(ctx, 5, 5)) | 357 | + assert.NilError(t, engine.feedback.AutoDismissStale(ctx, 5, 5)) |
| 357 | 358 | ||
| 358 | - dismissed, err := engine.IsFeedDismissed(ctx, "did:test:alice", "https://acted.com/feed") | 359 | + dismissed, err := engine.feedback.IsFeedDismissed(ctx, "did:test:alice", "https://acted.com/feed") |
| 359 | assert.NilError(t, err) | 360 | assert.NilError(t, err) |
| 360 | assert.Assert(t, !dismissed, "acted recommendation should not be auto-dismissed") | 361 | assert.Assert(t, !dismissed, "acted recommendation should not be auto-dismissed") |
| 361 | } | 362 | } |
| @@ -503,7 +504,7 @@ func TestDismissArticle(t *testing.T) { | |||
| 503 | 504 | ||
| 504 | engine := newTestEngine(dbs) | 505 | engine := newTestEngine(dbs) |
| 505 | 506 | ||
| 506 | - assert.NilError(t, engine.DismissArticle(ctx, "did:test:alice", "https://a.com/article1", "not_interested")) | 507 | + assert.NilError(t, engine.feedback.DismissArticle(ctx, "did:test:alice", "https://a.com/article1", "not_interested")) |
| 507 | 508 | ||
| 508 | var count int | 509 | var count int |
| 509 | assert.NilError(t, dbs.SQLDB().QueryRowContext(ctx, | 510 | assert.NilError(t, dbs.SQLDB().QueryRowContext(ctx, |
| @@ -531,8 +532,8 @@ func TestDismissFeed_Idempotent(t *testing.T) { | |||
| 531 | 532 | ||
| 532 | engine := newTestEngine(dbs) | 533 | engine := newTestEngine(dbs) |
| 533 | 534 | ||
| 534 | - assert.NilError(t, engine.DismissFeed(ctx, "did:test:alice", "https://a.com/feed", "reason1")) | 535 | + assert.NilError(t, engine.feedback.DismissFeed(ctx, "did:test:alice", "https://a.com/feed", "reason1")) |
| 535 | - assert.NilError(t, engine.DismissFeed(ctx, "did:test:alice", "https://a.com/feed", "reason2")) | 536 | + assert.NilError(t, engine.feedback.DismissFeed(ctx, "did:test:alice", "https://a.com/feed", "reason2")) |
| 536 | 537 | ||
| 537 | var count int | 538 | var count int |
| 538 | assert.NilError(t, dbs.SQLDB().QueryRowContext(ctx, | 539 | assert.NilError(t, dbs.SQLDB().QueryRowContext(ctx, |
modified
internal/cluster/scoring.go +2 -1 | @@ -7,6 +7,7 @@ import ( | ||
| 7 | 7 | "strings" |
| 8 | 8 | "time" |
| 9 | 9 | |
| 10 | + "pkg.rbrt.fr/glean/internal/ai" | |
| 10 | 11 | "pkg.rbrt.fr/glean/internal/db" |
| 11 | 12 | ) |
| 12 | 13 | |
| @@ -384,7 +385,7 @@ func (e *Engine) coldStartFromEmbeddings(ctx context.Context, userDID string, li | ||
| 384 | 385 | subSet[u] = true |
| 385 | 386 | } |
| 386 | 387 | |
| 387 | - queryBlob, err := avgEmbeddings(blobs, dim) | |
| 388 | + queryBlob, err := ai.AvgEmbeddings(blobs, dim) | |
| 388 | 389 | if err != nil { |
| 389 | 390 | return nil, fmt.Errorf("serialize query vector: %w", err) |
| 390 | 391 | } |
| @@ -7,6 +7,7 @@ import ( | |||
| 7 | "strings" | 7 | "strings" |
| 8 | "time" | 8 | "time" |
| 9 | 9 | ||
| 10 | + "pkg.rbrt.fr/glean/internal/ai" | ||
| 10 | "pkg.rbrt.fr/glean/internal/db" | 11 | "pkg.rbrt.fr/glean/internal/db" |
| 11 | ) | 12 | ) |
| 12 | 13 | ||
| @@ -384,7 +385,7 @@ func (e *Engine) coldStartFromEmbeddings(ctx context.Context, userDID string, li | |||
| 384 | subSet[u] = true | 385 | subSet[u] = true |
| 385 | } | 386 | } |
| 386 | 387 | ||
| 387 | - queryBlob, err := avgEmbeddings(blobs, dim) | 388 | + queryBlob, err := ai.AvgEmbeddings(blobs, dim) |
| 388 | if err != nil { | 389 | if err != nil { |
| 389 | return nil, fmt.Errorf("serialize query vector: %w", err) | 390 | return nil, fmt.Errorf("serialize query vector: %w", err) |
| 390 | } | 391 | } |
renamed
internal/feedback/feedback.go +26 -23 | similarity index 53% | ||
| rename from internal/cluster/dismiss.go | ||
| rename to internal/feedback/feedback.go | ||
| @@ -1,20 +1,26 @@ | ||
| 1 | -package cluster | |
| 1 | +package feedback | |
| 2 | 2 | |
| 3 | 3 | import ( |
| 4 | 4 | "context" |
| 5 | + "database/sql" | |
| 5 | 6 | "time" |
| 6 | 7 | ) |
| 7 | 8 | |
| 8 | -// Impression records that a recommendation was shown to a user. | |
| 9 | 9 | type Impression struct { |
| 10 | 10 | TargetType string |
| 11 | 11 | TargetID string |
| 12 | 12 | } |
| 13 | 13 | |
| 14 | -// Dismiss records that the user dismissed a recommendation. targetType is one | |
| 15 | -// of "feed", "article", "person". | |
| 16 | -func (e *Engine) Dismiss(ctx context.Context, userDID, targetType, targetID, reason string) error { | |
| 17 | - _, err := e.db.ExecContext(ctx, ` | |
| 14 | +type Service struct { | |
| 15 | + db *sql.DB | |
| 16 | +} | |
| 17 | + | |
| 18 | +func NewService(db *sql.DB) *Service { | |
| 19 | + return &Service{db: db} | |
| 20 | +} | |
| 21 | + | |
| 22 | +func (s *Service) Dismiss(ctx context.Context, userDID, targetType, targetID, reason string) error { | |
| 23 | + _, err := s.db.ExecContext(ctx, ` | |
| 18 | 24 | INSERT INTO main.dismissed_recommendations (user_did, target_type, target_id, reason) |
| 19 | 25 | VALUES (?, ?, ?, ?) |
| 20 | 26 | ON CONFLICT(user_did, target_type, target_id) DO UPDATE SET reason = excluded.reason, dismissed_at = CURRENT_TIMESTAMP |
| @@ -22,20 +28,20 @@ func (e *Engine) Dismiss(ctx context.Context, userDID, targetType, targetID, rea | ||
| 22 | 28 | return err |
| 23 | 29 | } |
| 24 | 30 | |
| 25 | -func (e *Engine) DismissFeed(ctx context.Context, userDID, feedURL, reason string) error { | |
| 26 | - return e.Dismiss(ctx, userDID, "feed", feedURL, reason) | |
| 31 | +func (s *Service) DismissFeed(ctx context.Context, userDID, feedURL, reason string) error { | |
| 32 | + return s.Dismiss(ctx, userDID, "feed", feedURL, reason) | |
| 27 | 33 | } |
| 28 | 34 | |
| 29 | -func (e *Engine) DismissArticle(ctx context.Context, userDID, articleURL, reason string) error { | |
| 30 | - return e.Dismiss(ctx, userDID, "article", articleURL, reason) | |
| 35 | +func (s *Service) DismissArticle(ctx context.Context, userDID, articleURL, reason string) error { | |
| 36 | + return s.Dismiss(ctx, userDID, "article", articleURL, reason) | |
| 31 | 37 | } |
| 32 | 38 | |
| 33 | -func (e *Engine) DismissPerson(ctx context.Context, userDID, targetDID, reason string) error { | |
| 34 | - return e.Dismiss(ctx, userDID, "person", targetDID, reason) | |
| 39 | +func (s *Service) DismissPerson(ctx context.Context, userDID, targetDID, reason string) error { | |
| 40 | + return s.Dismiss(ctx, userDID, "person", targetDID, reason) | |
| 35 | 41 | } |
| 36 | 42 | |
| 37 | -func (e *Engine) RecordImpressions(ctx context.Context, userDID string, impressions []Impression) error { | |
| 38 | - tx, err := e.db.BeginTx(ctx, nil) | |
| 43 | +func (s *Service) RecordImpressions(ctx context.Context, userDID string, impressions []Impression) error { | |
| 44 | + tx, err := s.db.BeginTx(ctx, nil) | |
| 39 | 45 | if err != nil { |
| 40 | 46 | return err |
| 41 | 47 | } |
| @@ -56,21 +62,18 @@ func (e *Engine) RecordImpressions(ctx context.Context, userDID string, impressi | ||
| 56 | 62 | return tx.Commit() |
| 57 | 63 | } |
| 58 | 64 | |
| 59 | -func (e *Engine) MarkImpressionActed(ctx context.Context, userDID, targetType, targetID string) error { | |
| 60 | - _, err := e.db.ExecContext(ctx, ` | |
| 65 | +func (s *Service) MarkImpressionActed(ctx context.Context, userDID, targetType, targetID string) error { | |
| 66 | + _, err := s.db.ExecContext(ctx, ` | |
| 61 | 67 | UPDATE main.recommendation_impressions SET acted = 1 |
| 62 | 68 | WHERE user_did = ? AND target_type = ? AND target_id = ? |
| 63 | 69 | `, userDID, targetType, targetID) |
| 64 | 70 | return err |
| 65 | 71 | } |
| 66 | 72 | |
| 67 | -// AutoDismissStale marks recommendations as dismissed if they were shown at | |
| 68 | -// least minShownCount times over more than maxAgeDays without the user acting | |
| 69 | -// on them. | |
| 70 | -func (e *Engine) AutoDismissStale(ctx context.Context, minShownCount int, maxAgeDays int) error { | |
| 73 | +func (s *Service) AutoDismissStale(ctx context.Context, minShownCount int, maxAgeDays int) error { | |
| 71 | 74 | cutoff := time.Now().AddDate(0, 0, -maxAgeDays).Format(time.RFC3339) |
| 72 | 75 | |
| 73 | - _, err := e.db.ExecContext(ctx, ` | |
| 76 | + _, err := s.db.ExecContext(ctx, ` | |
| 74 | 77 | INSERT OR IGNORE INTO main.dismissed_recommendations (user_did, target_type, target_id, reason, dismissed_at) |
| 75 | 78 | SELECT user_did, target_type, target_id, 'auto_stale', CURRENT_TIMESTAMP |
| 76 | 79 | FROM main.recommendation_impressions |
| @@ -81,9 +84,9 @@ func (e *Engine) AutoDismissStale(ctx context.Context, minShownCount int, maxAge | ||
| 81 | 84 | return err |
| 82 | 85 | } |
| 83 | 86 | |
| 84 | -func (e *Engine) IsFeedDismissed(ctx context.Context, userDID, feedURL string) (bool, error) { | |
| 87 | +func (s *Service) IsFeedDismissed(ctx context.Context, userDID, feedURL string) (bool, error) { | |
| 85 | 88 | var count int |
| 86 | - err := e.db.QueryRowContext(ctx, ` | |
| 89 | + err := s.db.QueryRowContext(ctx, ` | |
| 87 | 90 | SELECT COUNT(1) FROM main.dismissed_recommendations |
| 88 | 91 | WHERE user_did = ? AND target_type = 'feed' AND target_id = ? |
| 89 | 92 | `, userDID, feedURL).Scan(&count) |
| similarity index 53% | |||
| rename from internal/cluster/dismiss.go | |||
| rename to internal/feedback/feedback.go | |||
| @@ -1,20 +1,26 @@ | |||
| 1 | -package cluster | 1 | +package feedback |
| 2 | 2 | ||
| 3 | import ( | 3 | import ( |
| 4 | "context" | 4 | "context" |
| 5 | + "database/sql" | ||
| 5 | "time" | 6 | "time" |
| 6 | ) | 7 | ) |
| 7 | 8 | ||
| 8 | -// Impression records that a recommendation was shown to a user. | ||
| 9 | type Impression struct { | 9 | type Impression struct { |
| 10 | TargetType string | 10 | TargetType string |
| 11 | TargetID string | 11 | TargetID string |
| 12 | } | 12 | } |
| 13 | 13 | ||
| 14 | -// Dismiss records that the user dismissed a recommendation. targetType is one | 14 | +type Service struct { |
| 15 | -// of "feed", "article", "person". | 15 | + db *sql.DB |
| 16 | -func (e *Engine) Dismiss(ctx context.Context, userDID, targetType, targetID, reason string) error { | 16 | +} |
| 17 | - _, err := e.db.ExecContext(ctx, ` | 17 | + |
| 18 | +func NewService(db *sql.DB) *Service { | ||
| 19 | + return &Service{db: db} | ||
| 20 | +} | ||
| 21 | + | ||
| 22 | +func (s *Service) Dismiss(ctx context.Context, userDID, targetType, targetID, reason string) error { | ||
| 23 | + _, err := s.db.ExecContext(ctx, ` | ||
| 18 | INSERT INTO main.dismissed_recommendations (user_did, target_type, target_id, reason) | 24 | INSERT INTO main.dismissed_recommendations (user_did, target_type, target_id, reason) |
| 19 | VALUES (?, ?, ?, ?) | 25 | VALUES (?, ?, ?, ?) |
| 20 | ON CONFLICT(user_did, target_type, target_id) DO UPDATE SET reason = excluded.reason, dismissed_at = CURRENT_TIMESTAMP | 26 | ON CONFLICT(user_did, target_type, target_id) DO UPDATE SET reason = excluded.reason, dismissed_at = CURRENT_TIMESTAMP |
| @@ -22,20 +28,20 @@ func (e *Engine) Dismiss(ctx context.Context, userDID, targetType, targetID, rea | |||
| 22 | return err | 28 | return err |
| 23 | } | 29 | } |
| 24 | 30 | ||
| 25 | -func (e *Engine) DismissFeed(ctx context.Context, userDID, feedURL, reason string) error { | 31 | +func (s *Service) DismissFeed(ctx context.Context, userDID, feedURL, reason string) error { |
| 26 | - return e.Dismiss(ctx, userDID, "feed", feedURL, reason) | 32 | + return s.Dismiss(ctx, userDID, "feed", feedURL, reason) |
| 27 | } | 33 | } |
| 28 | 34 | ||
| 29 | -func (e *Engine) DismissArticle(ctx context.Context, userDID, articleURL, reason string) error { | 35 | +func (s *Service) DismissArticle(ctx context.Context, userDID, articleURL, reason string) error { |
| 30 | - return e.Dismiss(ctx, userDID, "article", articleURL, reason) | 36 | + return s.Dismiss(ctx, userDID, "article", articleURL, reason) |
| 31 | } | 37 | } |
| 32 | 38 | ||
| 33 | -func (e *Engine) DismissPerson(ctx context.Context, userDID, targetDID, reason string) error { | 39 | +func (s *Service) DismissPerson(ctx context.Context, userDID, targetDID, reason string) error { |
| 34 | - return e.Dismiss(ctx, userDID, "person", targetDID, reason) | 40 | + return s.Dismiss(ctx, userDID, "person", targetDID, reason) |
| 35 | } | 41 | } |
| 36 | 42 | ||
| 37 | -func (e *Engine) RecordImpressions(ctx context.Context, userDID string, impressions []Impression) error { | 43 | +func (s *Service) RecordImpressions(ctx context.Context, userDID string, impressions []Impression) error { |
| 38 | - tx, err := e.db.BeginTx(ctx, nil) | 44 | + tx, err := s.db.BeginTx(ctx, nil) |
| 39 | if err != nil { | 45 | if err != nil { |
| 40 | return err | 46 | return err |
| 41 | } | 47 | } |
| @@ -56,21 +62,18 @@ func (e *Engine) RecordImpressions(ctx context.Context, userDID string, impressi | |||
| 56 | return tx.Commit() | 62 | return tx.Commit() |
| 57 | } | 63 | } |
| 58 | 64 | ||
| 59 | -func (e *Engine) MarkImpressionActed(ctx context.Context, userDID, targetType, targetID string) error { | 65 | +func (s *Service) MarkImpressionActed(ctx context.Context, userDID, targetType, targetID string) error { |
| 60 | - _, err := e.db.ExecContext(ctx, ` | 66 | + _, err := s.db.ExecContext(ctx, ` |
| 61 | UPDATE main.recommendation_impressions SET acted = 1 | 67 | UPDATE main.recommendation_impressions SET acted = 1 |
| 62 | WHERE user_did = ? AND target_type = ? AND target_id = ? | 68 | WHERE user_did = ? AND target_type = ? AND target_id = ? |
| 63 | `, userDID, targetType, targetID) | 69 | `, userDID, targetType, targetID) |
| 64 | return err | 70 | return err |
| 65 | } | 71 | } |
| 66 | 72 | ||
| 67 | -// AutoDismissStale marks recommendations as dismissed if they were shown at | 73 | +func (s *Service) AutoDismissStale(ctx context.Context, minShownCount int, maxAgeDays int) error { |
| 68 | -// least minShownCount times over more than maxAgeDays without the user acting | ||
| 69 | -// on them. | ||
| 70 | -func (e *Engine) AutoDismissStale(ctx context.Context, minShownCount int, maxAgeDays int) error { | ||
| 71 | cutoff := time.Now().AddDate(0, 0, -maxAgeDays).Format(time.RFC3339) | 74 | cutoff := time.Now().AddDate(0, 0, -maxAgeDays).Format(time.RFC3339) |
| 72 | 75 | ||
| 73 | - _, err := e.db.ExecContext(ctx, ` | 76 | + _, err := s.db.ExecContext(ctx, ` |
| 74 | INSERT OR IGNORE INTO main.dismissed_recommendations (user_did, target_type, target_id, reason, dismissed_at) | 77 | INSERT OR IGNORE INTO main.dismissed_recommendations (user_did, target_type, target_id, reason, dismissed_at) |
| 75 | SELECT user_did, target_type, target_id, 'auto_stale', CURRENT_TIMESTAMP | 78 | SELECT user_did, target_type, target_id, 'auto_stale', CURRENT_TIMESTAMP |
| 76 | FROM main.recommendation_impressions | 79 | FROM main.recommendation_impressions |
| @@ -81,9 +84,9 @@ func (e *Engine) AutoDismissStale(ctx context.Context, minShownCount int, maxAge | |||
| 81 | return err | 84 | return err |
| 82 | } | 85 | } |
| 83 | 86 | ||
| 84 | -func (e *Engine) IsFeedDismissed(ctx context.Context, userDID, feedURL string) (bool, error) { | 87 | +func (s *Service) IsFeedDismissed(ctx context.Context, userDID, feedURL string) (bool, error) { |
| 85 | var count int | 88 | var count int |
| 86 | - err := e.db.QueryRowContext(ctx, ` | 89 | + err := s.db.QueryRowContext(ctx, ` |
| 87 | SELECT COUNT(1) FROM main.dismissed_recommendations | 90 | SELECT COUNT(1) FROM main.dismissed_recommendations |
| 88 | WHERE user_did = ? AND target_type = 'feed' AND target_id = ? | 91 | WHERE user_did = ? AND target_type = 'feed' AND target_id = ? |
| 89 | `, userDID, feedURL).Scan(&count) | 92 | `, userDID, feedURL).Scan(&count) |
modified
internal/server/articles_handler.go +2 -2 | @@ -373,7 +373,7 @@ func (s *Server) handleLikeArticle(w http.ResponseWriter, r *http.Request) { | ||
| 373 | 373 | http.Error(w, err.Error(), http.StatusInternalServerError) |
| 374 | 374 | return |
| 375 | 375 | } |
| 376 | - if err := s.engine.MarkImpressionActed(ctx, user.DID, "article", article.URL.String); err != nil { | |
| 376 | + if err := s.feedback.MarkImpressionActed(ctx, user.DID, "article", article.URL.String); err != nil { | |
| 377 | 377 | s.logger.Warn("failed to mark impression acted", "error", err) |
| 378 | 378 | } |
| 379 | 379 | sig := s.engine.GetDominantSignal(s.engine.GetWeights(ctx, user.DID)) |
| @@ -390,7 +390,7 @@ func (s *Server) handleLikeArticle(w http.ResponseWriter, r *http.Request) { | ||
| 390 | 390 | http.Error(w, err.Error(), http.StatusInternalServerError) |
| 391 | 391 | return |
| 392 | 392 | } |
| 393 | - if err := s.engine.MarkImpressionActed(ctx, user.DID, "article", article.URL.String); err != nil { | |
| 393 | + if err := s.feedback.MarkImpressionActed(ctx, user.DID, "article", article.URL.String); err != nil { | |
| 394 | 394 | s.logger.Warn("failed to mark impression acted", "error", err) |
| 395 | 395 | } |
| 396 | 396 | sig := s.engine.GetDominantSignal(s.engine.GetWeights(ctx, user.DID)) |
| @@ -373,7 +373,7 @@ func (s *Server) handleLikeArticle(w http.ResponseWriter, r *http.Request) { | |||
| 373 | http.Error(w, err.Error(), http.StatusInternalServerError) | 373 | http.Error(w, err.Error(), http.StatusInternalServerError) |
| 374 | return | 374 | return |
| 375 | } | 375 | } |
| 376 | - if err := s.engine.MarkImpressionActed(ctx, user.DID, "article", article.URL.String); err != nil { | 376 | + if err := s.feedback.MarkImpressionActed(ctx, user.DID, "article", article.URL.String); err != nil { |
| 377 | s.logger.Warn("failed to mark impression acted", "error", err) | 377 | s.logger.Warn("failed to mark impression acted", "error", err) |
| 378 | } | 378 | } |
| 379 | sig := s.engine.GetDominantSignal(s.engine.GetWeights(ctx, user.DID)) | 379 | sig := s.engine.GetDominantSignal(s.engine.GetWeights(ctx, user.DID)) |
| @@ -390,7 +390,7 @@ func (s *Server) handleLikeArticle(w http.ResponseWriter, r *http.Request) { | |||
| 390 | http.Error(w, err.Error(), http.StatusInternalServerError) | 390 | http.Error(w, err.Error(), http.StatusInternalServerError) |
| 391 | return | 391 | return |
| 392 | } | 392 | } |
| 393 | - if err := s.engine.MarkImpressionActed(ctx, user.DID, "article", article.URL.String); err != nil { | 393 | + if err := s.feedback.MarkImpressionActed(ctx, user.DID, "article", article.URL.String); err != nil { |
| 394 | s.logger.Warn("failed to mark impression acted", "error", err) | 394 | s.logger.Warn("failed to mark impression acted", "error", err) |
| 395 | } | 395 | } |
| 396 | sig := s.engine.GetDominantSignal(s.engine.GetWeights(ctx, user.DID)) | 396 | sig := s.engine.GetDominantSignal(s.engine.GetWeights(ctx, user.DID)) |
modified
internal/server/dashboard_handler.go +5 -4 | @@ -10,6 +10,7 @@ import ( | ||
| 10 | 10 | "pkg.rbrt.fr/glean/internal/atproto" |
| 11 | 11 | "pkg.rbrt.fr/glean/internal/cluster" |
| 12 | 12 | "pkg.rbrt.fr/glean/internal/db" |
| 13 | + "pkg.rbrt.fr/glean/internal/feedback" | |
| 13 | 14 | ) |
| 14 | 15 | |
| 15 | 16 | func (s *Server) handleDashboard(w http.ResponseWriter, r *http.Request) { |
| @@ -102,15 +103,15 @@ func (s *Server) handleDashboard(w http.ResponseWriter, r *http.Request) { | ||
| 102 | 103 | |
| 103 | 104 | resolvePeopleHandles(ctx, peopleRecs) |
| 104 | 105 | |
| 105 | - var impressions []cluster.Impression | |
| 106 | + var impressions []feedback.Impression | |
| 106 | 107 | for _, rec := range articleRecs { |
| 107 | - impressions = append(impressions, cluster.Impression{TargetType: "article", TargetID: rec.URL}) | |
| 108 | + impressions = append(impressions, feedback.Impression{TargetType: "article", TargetID: rec.URL}) | |
| 108 | 109 | } |
| 109 | 110 | for _, rec := range feedRecs { |
| 110 | - impressions = append(impressions, cluster.Impression{TargetType: "feed", TargetID: rec.FeedURL}) | |
| 111 | + impressions = append(impressions, feedback.Impression{TargetType: "feed", TargetID: rec.FeedURL}) | |
| 111 | 112 | } |
| 112 | 113 | if len(impressions) > 0 { |
| 113 | - if err := s.engine.RecordImpressions(ctx, user.DID, impressions); err != nil { | |
| 114 | + if err := s.feedback.RecordImpressions(ctx, user.DID, impressions); err != nil { | |
| 114 | 115 | s.logger.Warn("failed to record impressions", "error", err) |
| 115 | 116 | } |
| 116 | 117 | } |
| @@ -10,6 +10,7 @@ import ( | |||
| 10 | "pkg.rbrt.fr/glean/internal/atproto" | 10 | "pkg.rbrt.fr/glean/internal/atproto" |
| 11 | "pkg.rbrt.fr/glean/internal/cluster" | 11 | "pkg.rbrt.fr/glean/internal/cluster" |
| 12 | "pkg.rbrt.fr/glean/internal/db" | 12 | "pkg.rbrt.fr/glean/internal/db" |
| 13 | + "pkg.rbrt.fr/glean/internal/feedback" | ||
| 13 | ) | 14 | ) |
| 14 | 15 | ||
| 15 | func (s *Server) handleDashboard(w http.ResponseWriter, r *http.Request) { | 16 | func (s *Server) handleDashboard(w http.ResponseWriter, r *http.Request) { |
| @@ -102,15 +103,15 @@ func (s *Server) handleDashboard(w http.ResponseWriter, r *http.Request) { | |||
| 102 | 103 | ||
| 103 | resolvePeopleHandles(ctx, peopleRecs) | 104 | resolvePeopleHandles(ctx, peopleRecs) |
| 104 | 105 | ||
| 105 | - var impressions []cluster.Impression | 106 | + var impressions []feedback.Impression |
| 106 | for _, rec := range articleRecs { | 107 | for _, rec := range articleRecs { |
| 107 | - impressions = append(impressions, cluster.Impression{TargetType: "article", TargetID: rec.URL}) | 108 | + impressions = append(impressions, feedback.Impression{TargetType: "article", TargetID: rec.URL}) |
| 108 | } | 109 | } |
| 109 | for _, rec := range feedRecs { | 110 | for _, rec := range feedRecs { |
| 110 | - impressions = append(impressions, cluster.Impression{TargetType: "feed", TargetID: rec.FeedURL}) | 111 | + impressions = append(impressions, feedback.Impression{TargetType: "feed", TargetID: rec.FeedURL}) |
| 111 | } | 112 | } |
| 112 | if len(impressions) > 0 { | 113 | if len(impressions) > 0 { |
| 113 | - if err := s.engine.RecordImpressions(ctx, user.DID, impressions); err != nil { | 114 | + if err := s.feedback.RecordImpressions(ctx, user.DID, impressions); err != nil { |
| 114 | s.logger.Warn("failed to record impressions", "error", err) | 115 | s.logger.Warn("failed to record impressions", "error", err) |
| 115 | } | 116 | } |
| 116 | } | 117 | } |
modified
internal/server/feeds_handler.go +5 -4 | @@ -13,6 +13,7 @@ import ( | ||
| 13 | 13 | "pkg.rbrt.fr/glean/internal/cluster" |
| 14 | 14 | "pkg.rbrt.fr/glean/internal/db" |
| 15 | 15 | "pkg.rbrt.fr/glean/internal/feed" |
| 16 | + "pkg.rbrt.fr/glean/internal/feedback" | |
| 16 | 17 | ) |
| 17 | 18 | |
| 18 | 19 | func (s *Server) handleFeeds(w http.ResponseWriter, r *http.Request) { |
| @@ -106,11 +107,11 @@ func (s *Server) handleFeeds(w http.ResponseWriter, r *http.Request) { | ||
| 106 | 107 | |
| 107 | 108 | g2.Go(func() error { |
| 108 | 109 | if len(feedRecs) > 0 { |
| 109 | - impressions := make([]cluster.Impression, len(feedRecs)) | |
| 110 | + impressions := make([]feedback.Impression, len(feedRecs)) | |
| 110 | 111 | for i, rec := range feedRecs { |
| 111 | - impressions[i] = cluster.Impression{TargetType: "feed", TargetID: rec.FeedURL} | |
| 112 | + impressions[i] = feedback.Impression{TargetType: "feed", TargetID: rec.FeedURL} | |
| 112 | 113 | } |
| 113 | - if err := s.engine.RecordImpressions(gCtx2, user.DID, impressions); err != nil { | |
| 114 | + if err := s.feedback.RecordImpressions(gCtx2, user.DID, impressions); err != nil { | |
| 114 | 115 | s.logger.Warn("failed to record impressions", "error", err) |
| 115 | 116 | } |
| 116 | 117 | } |
| @@ -241,7 +242,7 @@ func (s *Server) handleAddFeed(w http.ResponseWriter, r *http.Request) { | ||
| 241 | 242 | |
| 242 | 243 | go s.storeFetchResult(context.WithoutCancel(r.Context()), feedURL, result.Feed.SiteURL, result) |
| 243 | 244 | |
| 244 | - if err := s.engine.MarkImpressionActed(r.Context(), user.DID, "feed", feedURL); err != nil { | |
| 245 | + if err := s.feedback.MarkImpressionActed(r.Context(), user.DID, "feed", feedURL); err != nil { | |
| 245 | 246 | s.logger.Warn("failed to mark impression acted", "error", err) |
| 246 | 247 | } |
| 247 | 248 | sig := s.engine.GetDominantSignal(s.engine.GetWeights(r.Context(), user.DID)) |
| @@ -13,6 +13,7 @@ import ( | |||
| 13 | "pkg.rbrt.fr/glean/internal/cluster" | 13 | "pkg.rbrt.fr/glean/internal/cluster" |
| 14 | "pkg.rbrt.fr/glean/internal/db" | 14 | "pkg.rbrt.fr/glean/internal/db" |
| 15 | "pkg.rbrt.fr/glean/internal/feed" | 15 | "pkg.rbrt.fr/glean/internal/feed" |
| 16 | + "pkg.rbrt.fr/glean/internal/feedback" | ||
| 16 | ) | 17 | ) |
| 17 | 18 | ||
| 18 | func (s *Server) handleFeeds(w http.ResponseWriter, r *http.Request) { | 19 | func (s *Server) handleFeeds(w http.ResponseWriter, r *http.Request) { |
| @@ -106,11 +107,11 @@ func (s *Server) handleFeeds(w http.ResponseWriter, r *http.Request) { | |||
| 106 | 107 | ||
| 107 | g2.Go(func() error { | 108 | g2.Go(func() error { |
| 108 | if len(feedRecs) > 0 { | 109 | if len(feedRecs) > 0 { |
| 109 | - impressions := make([]cluster.Impression, len(feedRecs)) | 110 | + impressions := make([]feedback.Impression, len(feedRecs)) |
| 110 | for i, rec := range feedRecs { | 111 | for i, rec := range feedRecs { |
| 111 | - impressions[i] = cluster.Impression{TargetType: "feed", TargetID: rec.FeedURL} | 112 | + impressions[i] = feedback.Impression{TargetType: "feed", TargetID: rec.FeedURL} |
| 112 | } | 113 | } |
| 113 | - if err := s.engine.RecordImpressions(gCtx2, user.DID, impressions); err != nil { | 114 | + if err := s.feedback.RecordImpressions(gCtx2, user.DID, impressions); err != nil { |
| 114 | s.logger.Warn("failed to record impressions", "error", err) | 115 | s.logger.Warn("failed to record impressions", "error", err) |
| 115 | } | 116 | } |
| 116 | } | 117 | } |
| @@ -241,7 +242,7 @@ func (s *Server) handleAddFeed(w http.ResponseWriter, r *http.Request) { | |||
| 241 | 242 | ||
| 242 | go s.storeFetchResult(context.WithoutCancel(r.Context()), feedURL, result.Feed.SiteURL, result) | 243 | go s.storeFetchResult(context.WithoutCancel(r.Context()), feedURL, result.Feed.SiteURL, result) |
| 243 | 244 | ||
| 244 | - if err := s.engine.MarkImpressionActed(r.Context(), user.DID, "feed", feedURL); err != nil { | 245 | + if err := s.feedback.MarkImpressionActed(r.Context(), user.DID, "feed", feedURL); err != nil { |
| 245 | s.logger.Warn("failed to mark impression acted", "error", err) | 246 | s.logger.Warn("failed to mark impression acted", "error", err) |
| 246 | } | 247 | } |
| 247 | sig := s.engine.GetDominantSignal(s.engine.GetWeights(r.Context(), user.DID)) | 248 | sig := s.engine.GetDominantSignal(s.engine.GetWeights(r.Context(), user.DID)) |
modified
internal/server/recs_handler.go +1 -1 | @@ -29,7 +29,7 @@ func (s *Server) handleDismiss(w http.ResponseWriter, r *http.Request, field, ta | ||
| 29 | 29 | reason = "not_interested" |
| 30 | 30 | } |
| 31 | 31 | |
| 32 | - if err := s.engine.Dismiss(r.Context(), user.DID, targetType, targetID, reason); err != nil { | |
| 32 | + if err := s.feedback.Dismiss(r.Context(), user.DID, targetType, targetID, reason); err != nil { | |
| 33 | 33 | s.logger.Error("failed to dismiss recommendation", "error", err, "type", targetType) |
| 34 | 34 | http.Error(w, err.Error(), http.StatusInternalServerError) |
| 35 | 35 | return |
| @@ -29,7 +29,7 @@ func (s *Server) handleDismiss(w http.ResponseWriter, r *http.Request, field, ta | |||
| 29 | reason = "not_interested" | 29 | reason = "not_interested" |
| 30 | } | 30 | } |
| 31 | 31 | ||
| 32 | - if err := s.engine.Dismiss(r.Context(), user.DID, targetType, targetID, reason); err != nil { | 32 | + if err := s.feedback.Dismiss(r.Context(), user.DID, targetType, targetID, reason); err != nil { |
| 33 | s.logger.Error("failed to dismiss recommendation", "error", err, "type", targetType) | 33 | s.logger.Error("failed to dismiss recommendation", "error", err, "type", targetType) |
| 34 | http.Error(w, err.Error(), http.StatusInternalServerError) | 34 | http.Error(w, err.Error(), http.StatusInternalServerError) |
| 35 | return | 35 | return |
modified
internal/server/server.go +3 -0 | @@ -26,6 +26,7 @@ import ( | ||
| 26 | 26 | "pkg.rbrt.fr/glean/internal/cluster" |
| 27 | 27 | "pkg.rbrt.fr/glean/internal/db" |
| 28 | 28 | "pkg.rbrt.fr/glean/internal/feed" |
| 29 | + "pkg.rbrt.fr/glean/internal/feedback" | |
| 29 | 30 | "pkg.rbrt.fr/glean/internal/metrics" |
| 30 | 31 | "pkg.rbrt.fr/glean/internal/sanitize" |
| 31 | 32 | "pkg.rbrt.fr/glean/internal/scraper" |
| @@ -67,6 +68,7 @@ type Server struct { | ||
| 67 | 68 | fetcher *feed.Fetcher |
| 68 | 69 | scheduler *feed.Scheduler |
| 69 | 70 | engine *cluster.Engine |
| 71 | + feedback *feedback.Service | |
| 70 | 72 | scraper *scraper.Scraper |
| 71 | 73 | clientID string |
| 72 | 74 | callbackURL string |
| @@ -98,6 +100,7 @@ func New(dbs *db.Store, clientID, callbackURL, addr string, scheduler *feed.Sche | ||
| 98 | 100 | fetcher: fetcher, |
| 99 | 101 | scheduler: scheduler, |
| 100 | 102 | engine: engine, |
| 103 | + feedback: feedback.NewService(dbs.SQLDB()), | |
| 101 | 104 | scraper: scraper.New(logger), |
| 102 | 105 | clientID: clientID, |
| 103 | 106 | callbackURL: callbackURL, |
| @@ -26,6 +26,7 @@ import ( | |||
| 26 | "pkg.rbrt.fr/glean/internal/cluster" | 26 | "pkg.rbrt.fr/glean/internal/cluster" |
| 27 | "pkg.rbrt.fr/glean/internal/db" | 27 | "pkg.rbrt.fr/glean/internal/db" |
| 28 | "pkg.rbrt.fr/glean/internal/feed" | 28 | "pkg.rbrt.fr/glean/internal/feed" |
| 29 | + "pkg.rbrt.fr/glean/internal/feedback" | ||
| 29 | "pkg.rbrt.fr/glean/internal/metrics" | 30 | "pkg.rbrt.fr/glean/internal/metrics" |
| 30 | "pkg.rbrt.fr/glean/internal/sanitize" | 31 | "pkg.rbrt.fr/glean/internal/sanitize" |
| 31 | "pkg.rbrt.fr/glean/internal/scraper" | 32 | "pkg.rbrt.fr/glean/internal/scraper" |
| @@ -67,6 +68,7 @@ type Server struct { | |||
| 67 | fetcher *feed.Fetcher | 68 | fetcher *feed.Fetcher |
| 68 | scheduler *feed.Scheduler | 69 | scheduler *feed.Scheduler |
| 69 | engine *cluster.Engine | 70 | engine *cluster.Engine |
| 71 | + feedback *feedback.Service | ||
| 70 | scraper *scraper.Scraper | 72 | scraper *scraper.Scraper |
| 71 | clientID string | 73 | clientID string |
| 72 | callbackURL string | 74 | callbackURL string |
| @@ -98,6 +100,7 @@ func New(dbs *db.Store, clientID, callbackURL, addr string, scheduler *feed.Sche | |||
| 98 | fetcher: fetcher, | 100 | fetcher: fetcher, |
| 99 | scheduler: scheduler, | 101 | scheduler: scheduler, |
| 100 | engine: engine, | 102 | engine: engine, |
| 103 | + feedback: feedback.NewService(dbs.SQLDB()), | ||
| 101 | scraper: scraper.New(logger), | 104 | scraper: scraper.New(logger), |
| 102 | clientID: clientID, | 105 | clientID: clientID, |
| 103 | callbackURL: callbackURL, | 106 | callbackURL: callbackURL, |
modified
main.go +7 -5 | @@ -13,10 +13,12 @@ import ( | ||
| 13 | 13 | "syscall" |
| 14 | 14 | "time" |
| 15 | 15 | |
| 16 | + "pkg.rbrt.fr/glean/internal/ai" | |
| 16 | 17 | "pkg.rbrt.fr/glean/internal/atproto" |
| 17 | 18 | "pkg.rbrt.fr/glean/internal/cluster" |
| 18 | 19 | "pkg.rbrt.fr/glean/internal/db" |
| 19 | 20 | "pkg.rbrt.fr/glean/internal/feed" |
| 21 | + "pkg.rbrt.fr/glean/internal/feedback" | |
| 20 | 22 | "pkg.rbrt.fr/glean/internal/server" |
| 21 | 23 | |
| 22 | 24 | vec "github.com/asg017/sqlite-vec-go-bindings/cgo" |
| @@ -73,9 +75,9 @@ func main() { | ||
| 73 | 75 | siteFetcher := atproto.NewStandardSiteFetcher(logger) |
| 74 | 76 | scheduler := feed.NewScheduler(storeAdapter, siteFetcher, logger, *fetchInterval, 30*time.Minute) |
| 75 | 77 | |
| 76 | - var embedder cluster.Embedder | |
| 78 | + var embedder ai.Embedder | |
| 77 | 79 | if embedURL := envOr("GLEAN_EMBED_BASE_URL", ""); embedURL != "" { |
| 78 | - embedder = cluster.NewEmbedderClient(cluster.EmbedderClientConfig{ | |
| 80 | + embedder = ai.NewEmbedder(ai.EmbedConfig{ | |
| 79 | 81 | BaseURL: embedURL, |
| 80 | 82 | APIKey: envOr("GLEAN_EMBED_API_KEY", ""), |
| 81 | 83 | Model: envOr("GLEAN_EMBED_MODEL", "text-embedding-3-small"), |
| @@ -90,16 +92,16 @@ func main() { | ||
| 90 | 92 | } |
| 91 | 93 | } |
| 92 | 94 | |
| 93 | - var llm *cluster.LLMClient | |
| 95 | + var llm ai.TextModel | |
| 94 | 96 | if llmURL := envOr("GLEAN_LLM_BASE_URL", ""); llmURL != "" { |
| 95 | - llm = cluster.NewLLMClient(cluster.LLMClientConfig{ | |
| 97 | + llm = ai.NewLLM(ai.LLMConfig{ | |
| 96 | 98 | BaseURL: llmURL, |
| 97 | 99 | APIKey: envOr("GLEAN_LLM_API_KEY", ""), |
| 98 | 100 | Model: envOr("GLEAN_LLM_MODEL", "gpt-4o-mini"), |
| 99 | 101 | }) |
| 100 | 102 | } |
| 101 | 103 | |
| 102 | - engine := cluster.NewEngine(dbs.SQLDB(), dbs.Articles, embedder, llm, logger, *clusterInterval, cluster.DefaultConfig()) | |
| 104 | + engine := cluster.NewEngine(dbs.SQLDB(), dbs.Articles, embedder, llm, feedback.NewService(dbs.SQLDB()), logger, *clusterInterval, cluster.DefaultConfig()) | |
| 103 | 105 | |
| 104 | 106 | fetcher := feed.NewFetcher(siteFetcher) |
| 105 | 107 | srv := server.New(dbs, clientID, callbackURL, *addr, scheduler, fetcher, engine, logger, []byte(sessionKey)) |
| @@ -13,10 +13,12 @@ import ( | |||
| 13 | "syscall" | 13 | "syscall" |
| 14 | "time" | 14 | "time" |
| 15 | 15 | ||
| 16 | + "pkg.rbrt.fr/glean/internal/ai" | ||
| 16 | "pkg.rbrt.fr/glean/internal/atproto" | 17 | "pkg.rbrt.fr/glean/internal/atproto" |
| 17 | "pkg.rbrt.fr/glean/internal/cluster" | 18 | "pkg.rbrt.fr/glean/internal/cluster" |
| 18 | "pkg.rbrt.fr/glean/internal/db" | 19 | "pkg.rbrt.fr/glean/internal/db" |
| 19 | "pkg.rbrt.fr/glean/internal/feed" | 20 | "pkg.rbrt.fr/glean/internal/feed" |
| 21 | + "pkg.rbrt.fr/glean/internal/feedback" | ||
| 20 | "pkg.rbrt.fr/glean/internal/server" | 22 | "pkg.rbrt.fr/glean/internal/server" |
| 21 | 23 | ||
| 22 | vec "github.com/asg017/sqlite-vec-go-bindings/cgo" | 24 | vec "github.com/asg017/sqlite-vec-go-bindings/cgo" |
| @@ -73,9 +75,9 @@ func main() { | |||
| 73 | siteFetcher := atproto.NewStandardSiteFetcher(logger) | 75 | siteFetcher := atproto.NewStandardSiteFetcher(logger) |
| 74 | scheduler := feed.NewScheduler(storeAdapter, siteFetcher, logger, *fetchInterval, 30*time.Minute) | 76 | scheduler := feed.NewScheduler(storeAdapter, siteFetcher, logger, *fetchInterval, 30*time.Minute) |
| 75 | 77 | ||
| 76 | - var embedder cluster.Embedder | 78 | + var embedder ai.Embedder |
| 77 | if embedURL := envOr("GLEAN_EMBED_BASE_URL", ""); embedURL != "" { | 79 | if embedURL := envOr("GLEAN_EMBED_BASE_URL", ""); embedURL != "" { |
| 78 | - embedder = cluster.NewEmbedderClient(cluster.EmbedderClientConfig{ | 80 | + embedder = ai.NewEmbedder(ai.EmbedConfig{ |
| 79 | BaseURL: embedURL, | 81 | BaseURL: embedURL, |
| 80 | APIKey: envOr("GLEAN_EMBED_API_KEY", ""), | 82 | APIKey: envOr("GLEAN_EMBED_API_KEY", ""), |
| 81 | Model: envOr("GLEAN_EMBED_MODEL", "text-embedding-3-small"), | 83 | Model: envOr("GLEAN_EMBED_MODEL", "text-embedding-3-small"), |
| @@ -90,16 +92,16 @@ func main() { | |||
| 90 | } | 92 | } |
| 91 | } | 93 | } |
| 92 | 94 | ||
| 93 | - var llm *cluster.LLMClient | 95 | + var llm ai.TextModel |
| 94 | if llmURL := envOr("GLEAN_LLM_BASE_URL", ""); llmURL != "" { | 96 | if llmURL := envOr("GLEAN_LLM_BASE_URL", ""); llmURL != "" { |
| 95 | - llm = cluster.NewLLMClient(cluster.LLMClientConfig{ | 97 | + llm = ai.NewLLM(ai.LLMConfig{ |
| 96 | BaseURL: llmURL, | 98 | BaseURL: llmURL, |
| 97 | APIKey: envOr("GLEAN_LLM_API_KEY", ""), | 99 | APIKey: envOr("GLEAN_LLM_API_KEY", ""), |
| 98 | Model: envOr("GLEAN_LLM_MODEL", "gpt-4o-mini"), | 100 | Model: envOr("GLEAN_LLM_MODEL", "gpt-4o-mini"), |
| 99 | }) | 101 | }) |
| 100 | } | 102 | } |
| 101 | 103 | ||
| 102 | - engine := cluster.NewEngine(dbs.SQLDB(), dbs.Articles, embedder, llm, logger, *clusterInterval, cluster.DefaultConfig()) | 104 | + engine := cluster.NewEngine(dbs.SQLDB(), dbs.Articles, embedder, llm, feedback.NewService(dbs.SQLDB()), logger, *clusterInterval, cluster.DefaultConfig()) |
| 103 | 105 | ||
| 104 | fetcher := feed.NewFetcher(siteFetcher) | 106 | fetcher := feed.NewFetcher(siteFetcher) |
| 105 | srv := server.New(dbs, clientID, callbackURL, *addr, scheduler, fetcher, engine, logger, []byte(sessionKey)) | 107 | srv := server.New(dbs, clientID, callbackURL, *addr, scheduler, fetcher, engine, logger, []byte(sessionKey)) |