diff --git a/cmd/labelmaker/api.go b/cmd/labelmaker/api.go index 39542bb..b7b9d68 100644 --- a/cmd/labelmaker/api.go +++ b/cmd/labelmaker/api.go @@ -5,6 +5,7 @@ import ( "fmt" "net/http" "slices" + "strconv" "strings" "time" @@ -31,7 +32,7 @@ func (srv *Server) RunAPI(bind string) { http.HandleFunc("GET /", HomeEndpoint) http.HandleFunc("POST /xrpc/at.nsid.cobalt.label.createLabel", svcAuth.Middleware(srv.CreateLabelEndpoint, true)) http.HandleFunc("POST /xrpc/at.nsid.cobalt.label.negateLabel", svcAuth.Middleware(srv.NegateLabelEndpoint, true)) - http.HandleFunc("POST /xrpc/at.nsid.cobalt.label.getLabels", srv.GetLabelsEndpoint) + http.HandleFunc("GET /xrpc/com.atproto.label.queryLabels", srv.QueryLabelsEndpoint) http.HandleFunc("GET /xrpc/com.atproto.label.subscribeLabels", srv.SubscribeLabelsEndpoint) if err := srv.httpServer.ListenAndServe(); err != nil && err != http.ErrServerClosed { @@ -154,36 +155,45 @@ func (srv *Server) mutateLabelEndpoint(w http.ResponseWriter, r *http.Request, n } } -type GetLabelsResponse struct { +type QueryLabelsResponse struct { + Cursor string `json:"cursor,omitempty"` Labels []labeling.Label `json:"labels"` } -func (srv *Server) GetLabelsEndpoint(w http.ResponseWriter, r *http.Request) { +func (srv *Server) QueryLabelsEndpoint(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") // parse query params params := r.URL.Query() - subjects := params["subject"] - if len(subjects) == 0 { - http.Error(w, fmt.Sprintf(`{"error": "BadRequest", "message": "missing subject parameter"}`), http.StatusBadRequest) - return - } - for _, subj := range subjects { - _, err := syntax.ParseURI(subj) + uriPatterns := params["uriPatterns"] + var sources []syntax.DID + for _, raw := range params["sources"] { + did, err := syntax.ParseDID(raw) if err != nil { - http.Error(w, fmt.Sprintf(`{"error": "BadRequest", "message": "invalid subject URI: %s"}`, subj), http.StatusBadRequest) + http.Error(w, fmt.Sprintf(`{"error": "BadRequest", "message": "invalid source DID: %s"}`, raw), http.StatusBadRequest) + return + } + sources = append(sources, did) + } + limit := 50 + rawLimit := params.Get("limit") + if rawLimit != "" { + limit, err := strconv.Atoi(rawLimit) + if err != nil || limit < 1 || limit > 250 { + http.Error(w, fmt.Sprintf(`{"error": "BadRequest", "message": "invalid limit"}`), http.StatusBadRequest) return } } + cursor := params.Get("limit") - labels, err := srv.Store.GetLabels(r.Context(), subjects, nil) + labels, nextCursor, err := srv.Store.QueryLabels(r.Context(), uriPatterns, sources, cursor, limit) if err != nil { http.Error(w, fmt.Sprintf(`{"error": "InternalServerError", "message": "%s"}`, err.Error()), http.StatusInternalServerError) return } - - respBody := GetLabelsResponse{ + respBody := QueryLabelsResponse{ Labels: labels, + Cursor: nextCursor, } if err := json.NewEncoder(w).Encode(respBody); err != nil { diff --git a/cmd/labelmaker/store.go b/cmd/labelmaker/store.go index 07af308..6e6965c 100644 --- a/cmd/labelmaker/store.go +++ b/cmd/labelmaker/store.go @@ -177,16 +177,32 @@ func (s *Store) LatestSeq(ctx context.Context) (int64, error) { return int64(row.ID), nil } -func (s *Store) QueryLabels(ctx context.Context, subjectPatterns []string, sources []syntax.DID, cursor string, limit int) ([]labeling.Label, error) { +func (s *Store) QueryLabels(ctx context.Context, subjectPatterns []string, sources []syntax.DID, cursor string, limit int) ([]labeling.Label, string, error) { var rows []Label + if limit <= 0 { + limit = 50 + } q := s.db.Where("neg = ?", false).Where( s.db.Where("exp IS NULL").Or("exp > ?", time.Now()), ).Order("id asc").Limit(limit) - // XXX: "LIKE" globbing instead of exact matches - q = q.Where("uri IN ?", subjectPatterns) + anyGlob := false + match := s.db + for _, uri := range subjectPatterns { + if strings.HasSuffix(uri, "*") { + anyGlob = true + match = match.Or("uri LIKE ?", strings.TrimSuffix(uri, "*")+"%") + } else { + match = match.Or("uri = ?", uri) + } + } + if anyGlob { + q = q.Where(match) + } else if len(subjectPatterns) > 0 { + q = q.Where("uri IN ?", subjectPatterns) + } if len(sources) > 0 { q = q.Where("src IN ?", sources) @@ -195,21 +211,28 @@ func (s *Store) QueryLabels(ctx context.Context, subjectPatterns []string, sourc if cursor != "" { seq, err := strconv.Atoi(cursor) if err != nil { - return nil, err + return nil, "", err } q = q.Where("id > ?", seq) } + q = q.Debug() + if err := q.Find(&rows).Error; err != nil { - return nil, err + return nil, "", err } out := make([]labeling.Label, len(rows)) for i, row := range rows { if err := json.Unmarshal(row.Data, &out[i]); err != nil { - return nil, err + return nil, "", err } } - return out, nil + nextCursor := "" + if len(rows) > 0 && len(rows) == limit { + nextCursor = fmt.Sprintf("%d", rows[len(rows)-1].ID) + } + + return out, nextCursor, nil } diff --git a/cmd/labelmaker/store_test.go b/cmd/labelmaker/store_test.go index 3e4fc1b..161f325 100644 --- a/cmd/labelmaker/store_test.go +++ b/cmd/labelmaker/store_test.go @@ -88,6 +88,38 @@ func TestStore(t *testing.T) { require.NoError(err) assert.Equal(int64(3), seq) + batch, _, err = s.QueryLabels(ctx, nil, nil, "", 50) + require.NoError(err) + assert.Equal(3, len(batch)) + + batch, _, err = s.QueryLabels(ctx, nil, nil, "999", 50) + require.NoError(err) + assert.Equal(0, len(batch)) + + batch, _, err = s.QueryLabels(ctx, nil, []syntax.DID{did1, did2}, "", 50) + require.NoError(err) + assert.Equal(3, len(batch)) + + batch, _, err = s.QueryLabels(ctx, nil, []syntax.DID{did2}, "", 50) + require.NoError(err) + assert.Equal(0, len(batch)) + + batch, _, err = s.QueryLabels(ctx, []string{"at://did:web:alice.example/app.bsky.feed.post/bbb"}, nil, "", 50) + require.NoError(err) + assert.Equal(2, len(batch)) + + batch, _, err = s.QueryLabels(ctx, []string{"at://did:web:alice.example/app.*"}, nil, "", 50) + require.NoError(err) + assert.Equal(3, len(batch)) + + batch, _, err = s.QueryLabels(ctx, []string{"*"}, nil, "", 50) + require.NoError(err) + assert.Equal(3, len(batch)) + + batch, _, err = s.QueryLabels(ctx, []string{"*/aaa"}, nil, "", 50) + require.NoError(err) + assert.Equal(0, len(batch)) + require.NoError(s.SaveLabels(ctx, []labeling.Label{l1neg})) seq, err = s.LatestSeq(ctx)