perf: improve follow distances computationUnverified
0ba006b parent: 7d7fc0d modified
internal/cluster/jaccard_test.go +3 -8 | @@ -284,7 +284,7 @@ func TestComputeFollowDistances(t *testing.T) { | ||
| 284 | 284 | assert.Assert(t, exists == 1, "alice should reach dave via 3 hops") |
| 285 | 285 | } |
| 286 | 286 | |
| 287 | -func TestComputeFollowDistancesData_SplitReadWrite(t *testing.T) { | |
| 287 | +func TestComputeFollowDistances_WritesPairsToDB(t *testing.T) { | |
| 288 | 288 | ctx := context.Background() |
| 289 | 289 | dbs := setupClusterTestDB(t) |
| 290 | 290 | seedClusterData(t, ctx, dbs) |
| @@ -292,16 +292,11 @@ func TestComputeFollowDistancesData_SplitReadWrite(t *testing.T) { | ||
| 292 | 292 | |
| 293 | 293 | engine := newTestEngine(dbs) |
| 294 | 294 | |
| 295 | - sources := []string{"did:test:alice", "did:test:bob", "did:test:carol", "did:test:dave"} | |
| 296 | - distances, err := engine.ComputeFollowDistancesData(ctx, sources) | |
| 297 | - assert.NilError(t, err) | |
| 298 | - assert.Assert(t, len(distances) > 0, "expected follow distance pairs") | |
| 299 | - | |
| 300 | - assert.NilError(t, engine.WriteFollowDistances(ctx, distances)) | |
| 295 | + assert.NilError(t, engine.ComputeFollowDistances(ctx)) | |
| 301 | 296 | |
| 302 | 297 | var count int |
| 303 | 298 | assert.NilError(t, dbs.SQLDB().QueryRowContext(ctx, `SELECT COUNT(*) FROM recs.follow_distances`).Scan(&count)) |
| 304 | - assert.Equal(t, count, len(distances)) | |
| 299 | + assert.Assert(t, count > 0, "expected follow distance pairs") | |
| 305 | 300 | } |
| 306 | 301 | |
| 307 | 302 | func TestAutoDismissStale(t *testing.T) { |
| @@ -284,7 +284,7 @@ func TestComputeFollowDistances(t *testing.T) { | |||
| 284 | assert.Assert(t, exists == 1, "alice should reach dave via 3 hops") | 284 | assert.Assert(t, exists == 1, "alice should reach dave via 3 hops") |
| 285 | } | 285 | } |
| 286 | 286 | ||
| 287 | -func TestComputeFollowDistancesData_SplitReadWrite(t *testing.T) { | 287 | +func TestComputeFollowDistances_WritesPairsToDB(t *testing.T) { |
| 288 | ctx := context.Background() | 288 | ctx := context.Background() |
| 289 | dbs := setupClusterTestDB(t) | 289 | dbs := setupClusterTestDB(t) |
| 290 | seedClusterData(t, ctx, dbs) | 290 | seedClusterData(t, ctx, dbs) |
| @@ -292,16 +292,11 @@ func TestComputeFollowDistancesData_SplitReadWrite(t *testing.T) { | |||
| 292 | 292 | ||
| 293 | engine := newTestEngine(dbs) | 293 | engine := newTestEngine(dbs) |
| 294 | 294 | ||
| 295 | - sources := []string{"did:test:alice", "did:test:bob", "did:test:carol", "did:test:dave"} | 295 | + assert.NilError(t, engine.ComputeFollowDistances(ctx)) |
| 296 | - distances, err := engine.ComputeFollowDistancesData(ctx, sources) | ||
| 297 | - assert.NilError(t, err) | ||
| 298 | - assert.Assert(t, len(distances) > 0, "expected follow distance pairs") | ||
| 299 | - | ||
| 300 | - assert.NilError(t, engine.WriteFollowDistances(ctx, distances)) | ||
| 301 | 296 | ||
| 302 | var count int | 297 | var count int |
| 303 | assert.NilError(t, dbs.SQLDB().QueryRowContext(ctx, `SELECT COUNT(*) FROM recs.follow_distances`).Scan(&count)) | 298 | assert.NilError(t, dbs.SQLDB().QueryRowContext(ctx, `SELECT COUNT(*) FROM recs.follow_distances`).Scan(&count)) |
| 304 | - assert.Equal(t, count, len(distances)) | 299 | + assert.Assert(t, count > 0, "expected follow distance pairs") |
| 305 | } | 300 | } |
| 306 | 301 | ||
| 307 | func TestAutoDismissStale(t *testing.T) { | 302 | func TestAutoDismissStale(t *testing.T) { |
modified
internal/cluster/social.go +21 -71 | @@ -2,6 +2,7 @@ package cluster | ||
| 2 | 2 | |
| 3 | 3 | import ( |
| 4 | 4 | "context" |
| 5 | + "database/sql" | |
| 5 | 6 | "fmt" |
| 6 | 7 | ) |
| 7 | 8 | |
| @@ -17,39 +18,25 @@ func chunk[T any](s []T, size int) [][]T { | ||
| 17 | 18 | return chunks |
| 18 | 19 | } |
| 19 | 20 | |
| 20 | -type followDistance struct { | |
| 21 | - userA string | |
| 22 | - userB string | |
| 23 | - distance int | |
| 24 | -} | |
| 25 | - | |
| 26 | -func (e *Engine) ComputeFollowDistancesData(ctx context.Context, sources []string) ([]followDistance, error) { | |
| 27 | - if len(sources) == 0 { | |
| 28 | - return nil, nil | |
| 29 | - } | |
| 30 | - | |
| 31 | - type pair struct { | |
| 32 | - src, dst string | |
| 21 | +func (e *Engine) writeFollowDistancesForUser(ctx context.Context, tx *sql.Tx, src string) (int, error) { | |
| 22 | + reachable, err := e.bfsReachable(ctx, src) | |
| 23 | + if err != nil { | |
| 24 | + return 0, err | |
| 33 | 25 | } |
| 34 | - distances := make(map[pair]int) | |
| 35 | 26 | |
| 36 | - for _, src := range sources { | |
| 37 | - reachable, err := e.bfsReachable(ctx, src) | |
| 38 | - if err != nil { | |
| 39 | - return nil, err | |
| 40 | - } | |
| 41 | - for other, d := range reachable { | |
| 42 | - if d > 0 { | |
| 43 | - distances[pair{src, other}] = d | |
| 27 | + written := 0 | |
| 28 | + for dst, d := range reachable { | |
| 29 | + if d > 0 { | |
| 30 | + if _, err := tx.ExecContext(ctx, | |
| 31 | + `INSERT INTO recs.follow_distances (user_a, user_b, distance) VALUES (?, ?, ?)`, | |
| 32 | + src, dst, d, | |
| 33 | + ); err != nil { | |
| 34 | + return written, err | |
| 44 | 35 | } |
| 36 | + written++ | |
| 45 | 37 | } |
| 46 | 38 | } |
| 47 | - | |
| 48 | - var result []followDistance | |
| 49 | - for k, d := range distances { | |
| 50 | - result = append(result, followDistance{userA: k.src, userB: k.dst, distance: d}) | |
| 51 | - } | |
| 52 | - return result, nil | |
| 39 | + return written, nil | |
| 53 | 40 | } |
| 54 | 41 | |
| 55 | 42 | func (e *Engine) bfsReachable(ctx context.Context, src string) (map[string]int, error) { |
| @@ -96,35 +83,6 @@ func (e *Engine) bfsReachable(ctx context.Context, src string) (map[string]int, | ||
| 96 | 83 | return reachable, nil |
| 97 | 84 | } |
| 98 | 85 | |
| 99 | -func (e *Engine) WriteFollowDistances(ctx context.Context, distances []followDistance) error { | |
| 100 | - tx, err := e.db.BeginTx(ctx, nil) | |
| 101 | - if err != nil { | |
| 102 | - return err | |
| 103 | - } | |
| 104 | - defer func() { _ = tx.Rollback() }() | |
| 105 | - | |
| 106 | - if _, err := tx.ExecContext(ctx, `DELETE FROM recs.follow_distances`); err != nil { | |
| 107 | - return err | |
| 108 | - } | |
| 109 | - | |
| 110 | - stmt, err := tx.PrepareContext(ctx, `INSERT INTO recs.follow_distances (user_a, user_b, distance) VALUES (?, ?, ?)`) | |
| 111 | - if err != nil { | |
| 112 | - return err | |
| 113 | - } | |
| 114 | - defer stmt.Close() | |
| 115 | - | |
| 116 | - for _, d := range distances { | |
| 117 | - if _, err := stmt.ExecContext(ctx, d.userA, d.userB, d.distance); err != nil { | |
| 118 | - return err | |
| 119 | - } | |
| 120 | - } | |
| 121 | - | |
| 122 | - e.logger.Info("follow distances computed", "pairs", len(distances)) | |
| 123 | - return tx.Commit() | |
| 124 | -} | |
| 125 | - | |
| 126 | -// ComputeFollowDistances incrementally recomputes follow distances for users | |
| 127 | -// whose follows changed since the last run, as tracked by the follows_dirty column. | |
| 128 | 86 | func (e *Engine) ComputeFollowDistances(ctx context.Context) error { |
| 129 | 87 | rows, err := e.db.QueryContext(ctx, `SELECT did FROM main.users WHERE follows_dirty = 1`) |
| 130 | 88 | if err != nil { |
| @@ -146,11 +104,6 @@ func (e *Engine) ComputeFollowDistances(ctx context.Context) error { | ||
| 146 | 104 | return nil |
| 147 | 105 | } |
| 148 | 106 | |
| 149 | - distances, err := e.ComputeFollowDistancesData(ctx, dirtyUsers) | |
| 150 | - if err != nil { | |
| 151 | - return err | |
| 152 | - } | |
| 153 | - | |
| 154 | 107 | tx, err := e.db.BeginTx(ctx, nil) |
| 155 | 108 | if err != nil { |
| 156 | 109 | return err |
| @@ -173,16 +126,13 @@ func (e *Engine) ComputeFollowDistances(ctx context.Context) error { | ||
| 173 | 126 | } |
| 174 | 127 | } |
| 175 | 128 | |
| 176 | - stmt, err := tx.PrepareContext(ctx, `INSERT INTO recs.follow_distances (user_a, user_b, distance) VALUES (?, ?, ?)`) | |
| 177 | - if err != nil { | |
| 178 | - return err | |
| 179 | - } | |
| 180 | - defer stmt.Close() | |
| 181 | - | |
| 182 | - for _, d := range distances { | |
| 183 | - if _, err := stmt.ExecContext(ctx, d.userA, d.userB, d.distance); err != nil { | |
| 129 | + var totalPairs int | |
| 130 | + for _, did := range dirtyUsers { | |
| 131 | + n, err := e.writeFollowDistancesForUser(ctx, tx, did) | |
| 132 | + if err != nil { | |
| 184 | 133 | return err |
| 185 | 134 | } |
| 135 | + totalPairs += n | |
| 186 | 136 | } |
| 187 | 137 | |
| 188 | 138 | for _, chunk := range chunk(dirtyUsers, sqliteMaxVars) { |
| @@ -200,6 +150,6 @@ func (e *Engine) ComputeFollowDistances(ctx context.Context) error { | ||
| 200 | 150 | } |
| 201 | 151 | } |
| 202 | 152 | |
| 203 | - e.logger.Info("follow distances computed", "users", len(dirtyUsers), "pairs", len(distances)) | |
| 153 | + e.logger.Info("follow distances computed", "users", len(dirtyUsers), "pairs", totalPairs) | |
| 204 | 154 | return tx.Commit() |
| 205 | 155 | } |
| @@ -2,6 +2,7 @@ package cluster | |||
| 2 | 2 | ||
| 3 | import ( | 3 | import ( |
| 4 | "context" | 4 | "context" |
| 5 | + "database/sql" | ||
| 5 | "fmt" | 6 | "fmt" |
| 6 | ) | 7 | ) |
| 7 | 8 | ||
| @@ -17,39 +18,25 @@ func chunk[T any](s []T, size int) [][]T { | |||
| 17 | return chunks | 18 | return chunks |
| 18 | } | 19 | } |
| 19 | 20 | ||
| 20 | -type followDistance struct { | 21 | +func (e *Engine) writeFollowDistancesForUser(ctx context.Context, tx *sql.Tx, src string) (int, error) { |
| 21 | - userA string | 22 | + reachable, err := e.bfsReachable(ctx, src) |
| 22 | - userB string | 23 | + if err != nil { |
| 23 | - distance int | 24 | + return 0, err |
| 24 | -} | ||
| 25 | - | ||
| 26 | -func (e *Engine) ComputeFollowDistancesData(ctx context.Context, sources []string) ([]followDistance, error) { | ||
| 27 | - if len(sources) == 0 { | ||
| 28 | - return nil, nil | ||
| 29 | - } | ||
| 30 | - | ||
| 31 | - type pair struct { | ||
| 32 | - src, dst string | ||
| 33 | } | 25 | } |
| 34 | - distances := make(map[pair]int) | ||
| 35 | 26 | ||
| 36 | - for _, src := range sources { | 27 | + written := 0 |
| 37 | - reachable, err := e.bfsReachable(ctx, src) | 28 | + for dst, d := range reachable { |
| 38 | - if err != nil { | 29 | + if d > 0 { |
| 39 | - return nil, err | 30 | + if _, err := tx.ExecContext(ctx, |
| 40 | - } | 31 | + `INSERT INTO recs.follow_distances (user_a, user_b, distance) VALUES (?, ?, ?)`, |
| 41 | - for other, d := range reachable { | 32 | + src, dst, d, |
| 42 | - if d > 0 { | 33 | + ); err != nil { |
| 43 | - distances[pair{src, other}] = d | 34 | + return written, err |
| 44 | } | 35 | } |
| 36 | + written++ | ||
| 45 | } | 37 | } |
| 46 | } | 38 | } |
| 47 | - | 39 | + return written, nil |
| 48 | - var result []followDistance | ||
| 49 | - for k, d := range distances { | ||
| 50 | - result = append(result, followDistance{userA: k.src, userB: k.dst, distance: d}) | ||
| 51 | - } | ||
| 52 | - return result, nil | ||
| 53 | } | 40 | } |
| 54 | 41 | ||
| 55 | func (e *Engine) bfsReachable(ctx context.Context, src string) (map[string]int, error) { | 42 | func (e *Engine) bfsReachable(ctx context.Context, src string) (map[string]int, error) { |
| @@ -96,35 +83,6 @@ func (e *Engine) bfsReachable(ctx context.Context, src string) (map[string]int, | |||
| 96 | return reachable, nil | 83 | return reachable, nil |
| 97 | } | 84 | } |
| 98 | 85 | ||
| 99 | -func (e *Engine) WriteFollowDistances(ctx context.Context, distances []followDistance) error { | ||
| 100 | - tx, err := e.db.BeginTx(ctx, nil) | ||
| 101 | - if err != nil { | ||
| 102 | - return err | ||
| 103 | - } | ||
| 104 | - defer func() { _ = tx.Rollback() }() | ||
| 105 | - | ||
| 106 | - if _, err := tx.ExecContext(ctx, `DELETE FROM recs.follow_distances`); err != nil { | ||
| 107 | - return err | ||
| 108 | - } | ||
| 109 | - | ||
| 110 | - stmt, err := tx.PrepareContext(ctx, `INSERT INTO recs.follow_distances (user_a, user_b, distance) VALUES (?, ?, ?)`) | ||
| 111 | - if err != nil { | ||
| 112 | - return err | ||
| 113 | - } | ||
| 114 | - defer stmt.Close() | ||
| 115 | - | ||
| 116 | - for _, d := range distances { | ||
| 117 | - if _, err := stmt.ExecContext(ctx, d.userA, d.userB, d.distance); err != nil { | ||
| 118 | - return err | ||
| 119 | - } | ||
| 120 | - } | ||
| 121 | - | ||
| 122 | - e.logger.Info("follow distances computed", "pairs", len(distances)) | ||
| 123 | - return tx.Commit() | ||
| 124 | -} | ||
| 125 | - | ||
| 126 | -// ComputeFollowDistances incrementally recomputes follow distances for users | ||
| 127 | -// whose follows changed since the last run, as tracked by the follows_dirty column. | ||
| 128 | func (e *Engine) ComputeFollowDistances(ctx context.Context) error { | 86 | func (e *Engine) ComputeFollowDistances(ctx context.Context) error { |
| 129 | rows, err := e.db.QueryContext(ctx, `SELECT did FROM main.users WHERE follows_dirty = 1`) | 87 | rows, err := e.db.QueryContext(ctx, `SELECT did FROM main.users WHERE follows_dirty = 1`) |
| 130 | if err != nil { | 88 | if err != nil { |
| @@ -146,11 +104,6 @@ func (e *Engine) ComputeFollowDistances(ctx context.Context) error { | |||
| 146 | return nil | 104 | return nil |
| 147 | } | 105 | } |
| 148 | 106 | ||
| 149 | - distances, err := e.ComputeFollowDistancesData(ctx, dirtyUsers) | ||
| 150 | - if err != nil { | ||
| 151 | - return err | ||
| 152 | - } | ||
| 153 | - | ||
| 154 | tx, err := e.db.BeginTx(ctx, nil) | 107 | tx, err := e.db.BeginTx(ctx, nil) |
| 155 | if err != nil { | 108 | if err != nil { |
| 156 | return err | 109 | return err |
| @@ -173,16 +126,13 @@ func (e *Engine) ComputeFollowDistances(ctx context.Context) error { | |||
| 173 | } | 126 | } |
| 174 | } | 127 | } |
| 175 | 128 | ||
| 176 | - stmt, err := tx.PrepareContext(ctx, `INSERT INTO recs.follow_distances (user_a, user_b, distance) VALUES (?, ?, ?)`) | 129 | + var totalPairs int |
| 177 | - if err != nil { | 130 | + for _, did := range dirtyUsers { |
| 178 | - return err | 131 | + n, err := e.writeFollowDistancesForUser(ctx, tx, did) |
| 179 | - } | 132 | + if err != nil { |
| 180 | - defer stmt.Close() | ||
| 181 | - | ||
| 182 | - for _, d := range distances { | ||
| 183 | - if _, err := stmt.ExecContext(ctx, d.userA, d.userB, d.distance); err != nil { | ||
| 184 | return err | 133 | return err |
| 185 | } | 134 | } |
| 135 | + totalPairs += n | ||
| 186 | } | 136 | } |
| 187 | 137 | ||
| 188 | for _, chunk := range chunk(dirtyUsers, sqliteMaxVars) { | 138 | for _, chunk := range chunk(dirtyUsers, sqliteMaxVars) { |
| @@ -200,6 +150,6 @@ func (e *Engine) ComputeFollowDistances(ctx context.Context) error { | |||
| 200 | } | 150 | } |
| 201 | } | 151 | } |
| 202 | 152 | ||
| 203 | - e.logger.Info("follow distances computed", "users", len(dirtyUsers), "pairs", len(distances)) | 153 | + e.logger.Info("follow distances computed", "users", len(dirtyUsers), "pairs", totalPairs) |
| 204 | return tx.Commit() | 154 | return tx.Commit() |
| 205 | } | 155 | } |