diff --git a/features/api/connections.go b/features/api/connections.go index 31e54dd..7fcc56f 100644 --- a/features/api/connections.go +++ b/features/api/connections.go @@ -81,7 +81,7 @@ func (h *Handlers) CreateConnection(w http.ResponseWriter, r *http.Request) { EventURI: r.URL.Query().Get("event"), } - uri, _, err := connection.Put(r.Context(), sess, rec) + uri, _, err := connection.Put(r.Context(), sess, h.DB, rec) if err != nil { slog.Warn("api: create connection", "from", did, "to", targetDID, "err", err) writeError(w, http.StatusInternalServerError, "failed to create connection") @@ -128,7 +128,7 @@ func (h *Handlers) FlushConnections(w http.ResponseWriter, r *http.Request) { continue } rec := connection.Record{With: target} - if _, _, err := connection.Put(r.Context(), sess, rec); err != nil { + if _, _, err := connection.Put(r.Context(), sess, h.DB, rec); err != nil { slog.Warn("api: flush connection", "target", target, "err", err) errs++ continue diff --git a/features/auth/link.go b/features/auth/link.go index ed6356e..8c12ed4 100644 --- a/features/auth/link.go +++ b/features/auth/link.go @@ -62,7 +62,7 @@ func (h *Handlers) LinkLocalToATProto(w http.ResponseWriter, r *http.Request) { EventURI: eventURI, ConnectedAt: time.Now().UTC(), } - if _, _, err := connection.Put(ctx, sess, rec); err != nil { + if _, _, err := connection.Put(ctx, sess, h.DB, rec); err != nil { slog.Warn("link: migrate connection", "target", targetDID, "err", err) } } diff --git a/features/connect/handlers.go b/features/connect/handlers.go index fb8b4dd..39cfe63 100644 --- a/features/connect/handlers.go +++ b/features/connect/handlers.go @@ -124,7 +124,7 @@ func (h *Handlers) Connect(w http.ResponseWriter, r *http.Request) { viewerPDS := h.lookupPDSForDID(r, viewerDID) if connection.HasConnection(r.Context(), viewerPDS, viewerDID, target, eventURI) { slog.Debug("connect: duplicate skipped", "viewer", viewerDID.String(), "target", target.String(), "event", eventURI) - } else if _, _, err := connection.Put(r.Context(), viewerSess, connRec); err != nil { + } else if _, _, err := connection.Put(r.Context(), viewerSess, h.DB, connRec); err != nil { slog.Warn("connect: viewer record", "viewer", viewerDID.String(), "target", target.String(), "err", err) } @@ -318,7 +318,7 @@ func (h *Handlers) ConnectFlushLocal(w http.ResponseWriter, r *http.Request) { skipped++ continue } - if _, _, err := connection.Put(r.Context(), viewerSess, connection.Record{With: target}); err != nil { + if _, _, err := connection.Put(r.Context(), viewerSess, h.DB, connection.Record{With: target}); err != nil { slog.Warn("flush-local: viewer record", "target", target.String(), "err", err) skipped++ errs = append(errs, "write failed for "+target.String()) diff --git a/internal/connection/connection.go b/internal/connection/connection.go index d549f43..c9f6261 100644 --- a/internal/connection/connection.go +++ b/internal/connection/connection.go @@ -26,6 +26,7 @@ package connection import ( "context" + "database/sql" "errors" "fmt" "time" @@ -59,7 +60,9 @@ type Record struct { // records — that's by design per the spec (e.g. across multiple events). // // Caller must have the `repo:quest.atmo.connection` scope on the session. -func Put(ctx context.Context, sess *oauth.ClientSession, rec Record) (uri, cid string, err error) { +// If db is non-nil, also writes a connection bookmark row (empty note) so +// CountConnectionsForDID can accurately count the user's connections. +func Put(ctx context.Context, sess *oauth.ClientSession, db *sql.DB, rec Record) (uri, cid string, err error) { if sess == nil { return "", "", errors.New("connection: nil oauth session") } @@ -85,10 +88,6 @@ func Put(ctx context.Context, sess *oauth.ClientSession, rec Record) (uri, cid s input := map[string]any{ "repo": sess.Data.AccountDID.String(), "collection": NSID, - // `validate` is intentionally omitted; PDSes that don't yet know - // the quest.atmo.connection lexicon reject `validate: true` with - // `InvalidRequest: Unknown lexicon type`. Leaving it unset asks - // the PDS to validate only against lexicons it already knows. "record": value, } var out struct { @@ -98,5 +97,16 @@ func Put(ctx context.Context, sess *oauth.ClientSession, rec Record) (uri, cid s if err := sess.APIClient().Post(ctx, syntax.NSID(nsidCreateRecord), input, &out); err != nil { return "", "", fmt.Errorf("createRecord %s: %w", NSID, err) } + + // Write connection bookmark (best-effort) + if db != nil { + viewerDID := sess.Data.AccountDID.String() + _, _ = db.ExecContext(ctx, ` + INSERT INTO connection_notes (viewer_did, target_did, notes, follow_up, updated_at) + VALUES (?, ?, '', 0, CURRENT_TIMESTAMP) + ON CONFLICT(viewer_did, target_did) DO NOTHING + `, viewerDID, rec.With.String()) + } + return out.URI, out.CID, nil } diff --git a/internal/connection/drain.go b/internal/connection/drain.go index d2cb27e..208e752 100644 --- a/internal/connection/drain.go +++ b/internal/connection/drain.go @@ -58,6 +58,15 @@ func Drain(ctx context.Context, q *Queue, sess *oauth.ClientSession, pdsHost str "event", item.EventURI, ) } + // Write bookmark for this known connection + if q.db != nil { + viewerDID := target.String() + _, _ = q.db.ExecContext(ctx, ` + INSERT INTO connection_notes (viewer_did, target_did, notes, follow_up, updated_at) + VALUES (?, ?, '', 0, CURRENT_TIMESTAMP) + ON CONFLICT(viewer_did, target_did) DO NOTHING + `, viewerDID, item.InitiatorDID.String()) + } // Delete the queue row — no need to retry. _ = q.Delete(ctx, item.ID) res.Skipped++ @@ -67,7 +76,7 @@ func Drain(ctx context.Context, q *Queue, sess *oauth.ClientSession, pdsHost str } continue } - _, _, err := Put(ctx, sess, Record{With: item.InitiatorDID, EventURI: item.EventURI}) + _, _, err := Put(ctx, sess, q.db, Record{With: item.InitiatorDID, EventURI: item.EventURI}) if err != nil { if logger != nil { logger.Warn("connection drain: write failed", diff --git a/internal/connection/local.go b/internal/connection/local.go index 21d42ce..da2716c 100644 --- a/internal/connection/local.go +++ b/internal/connection/local.go @@ -45,5 +45,24 @@ func WriteLocal(ctx context.Context, db *sql.DB, (viewer_did, viewer_local_id, target_did, target_local_id, event_uri, connected_at) VALUES (?, ?, ?, ?, ?, CURRENT_TIMESTAMP) `, vDID, vLocalID, tDID, tLocalID, eventURI) - return err + if err != nil { + return err + } + + // Write connection bookmark (best-effort) + viewerID := viewerDID + if viewerID == "" { + viewerID = viewerLocalID + } + targetID := targetDID + if targetID == "" { + targetID = targetLocalID + } + _, _ = db.ExecContext(ctx, ` + INSERT INTO connection_notes (viewer_did, target_did, notes, follow_up, updated_at) + VALUES (?, ?, '', 0, CURRENT_TIMESTAMP) + ON CONFLICT(viewer_did, target_did) DO NOTHING + `, viewerID, targetID) + + return nil } diff --git a/internal/notes/notes.go b/internal/notes/notes.go index 6f021d8..93e41ef 100644 --- a/internal/notes/notes.go +++ b/internal/notes/notes.go @@ -33,17 +33,10 @@ func Get(ctx context.Context, db *sql.DB, viewerDID, targetDID string) (Note, er return n, nil } -// Put upserts the note for (viewer, target). Empty notes with no follow-up -// flag deletes the row to keep the table tidy. +// Put upserts the note for (viewer, target). Always upserts — empty notes +// with no follow-up are kept as bookmarks so CountConnectionsForDID can +// accurately count a user's connections. func Put(ctx context.Context, db *sql.DB, viewerDID, targetDID string, n Note) error { - if n.Notes == "" && !n.FollowUp { - _, err := db.ExecContext(ctx, ` - DELETE FROM connection_notes - WHERE viewer_did = ? AND target_did = ? - `, viewerDID, targetDID) - return err - } - followUp := 0 if n.FollowUp { followUp = 1 diff --git a/internal/notes/notes_test.go b/internal/notes/notes_test.go index 18f07c9..818bbe9 100644 --- a/internal/notes/notes_test.go +++ b/internal/notes/notes_test.go @@ -78,19 +78,20 @@ func TestPut_UpsertOverwrites(t *testing.T) { } } -func TestPut_EmptyDeletesRow(t *testing.T) { +func TestPut_EmptyRowPersists(t *testing.T) { ctx, db := newTestDB(t) // Seed a note. if err := Put(ctx, db, viewerDID, targetDID, Note{Notes: "hi"}); err != nil { t.Fatalf("Put: %v", err) } - // Clear it — empty notes + no follow-up should delete the row. + // Clear it — empty notes + no follow-up should upsert (keep the row). if err := Put(ctx, db, viewerDID, targetDID, Note{}); err != nil { t.Fatalf("Put empty: %v", err) } + // Row should persist (empty, but still counted for connection bookmarks). got, _ := Get(ctx, db, viewerDID, targetDID) - if got.Notes != "" || got.FollowUp { - t.Errorf("after clearing: %+v; want zero Note", got) + if got.Notes != "" { + t.Errorf("Notes = %q; want empty", got.Notes) } }