diff --git a/README.md b/README.md index 8f93b61..5cfe1ce 100644 --- a/README.md +++ b/README.md @@ -256,7 +256,7 @@ Just because something is implemented doesn't mean it is finished. Tons of these ### Other -- [ ] `com.atproto.label.queryLabels` +- [x] `com.atproto.label.queryLabels` - [x] `com.atproto.moderation.createReport` (Note: this should be handled by proxying, not actually implemented in the PDS) - [x] `app.bsky.actor.getPreferences` - [x] `app.bsky.actor.putPreferences` diff --git a/cmd/cocoon/main.go b/cmd/cocoon/main.go index 410bb87..c1b7630 100644 --- a/cmd/cocoon/main.go +++ b/cmd/cocoon/main.go @@ -79,6 +79,11 @@ func main() { Name: "admin-password", EnvVars: []string{"COCOON_ADMIN_PASSWORD"}, }, + &cli.BoolFlag{ + Name: "require-invite", + EnvVars: []string{"COCOON_REQUIRE_INVITE"}, + Value: true, + }, &cli.StringFlag{ Name: "smtp-user", EnvVars: []string{"COCOON_SMTP_USER"}, @@ -185,6 +190,7 @@ var runServe = &cli.Command{ Version: Version, Relays: cmd.StringSlice("relays"), AdminPassword: cmd.String("admin-password"), + RequireInvite: cmd.Bool("require-invite"), SmtpUser: cmd.String("smtp-user"), SmtpPass: cmd.String("smtp-pass"), SmtpHost: cmd.String("smtp-host"), diff --git a/server/handle_label_query_labels.go b/server/handle_label_query_labels.go new file mode 100644 index 0000000..8e1c1d9 --- /dev/null +++ b/server/handle_label_query_labels.go @@ -0,0 +1,34 @@ +package server + +import ( + "github.com/labstack/echo/v4" +) + +type Label struct { + Ver *int `json:"ver,omitempty"` + Src string `json:"src"` + Uri string `json:"uri"` + Cid *string `json:"cid,omitempty"` + Val string `json:"val"` + Neg *bool `json:"neg,omitempty"` + Cts string `json:"cts"` + Exp *string `json:"exp,omitempty"` + Sig []byte `json:"sig,omitempty"` +} + +type ComAtprotoLabelQueryLabelsResponse struct { + Cursor *string `json:"cursor,omitempty"` + Labels []Label `json:"labels"` +} + +func (s *Server) handleLabelQueryLabels(e echo.Context) error { + svc := e.Request().Header.Get("atproto-proxy") + if svc != "" || s.config.FallbackProxy != "" { + return s.handleProxy(e) + } + + return e.JSON(200, ComAtprotoLabelQueryLabelsResponse{ + Cursor: nil, + Labels: []Label{}, + }) +} diff --git a/server/handle_server_create_account.go b/server/handle_server_create_account.go index 21f3aa6..b534acc 100644 --- a/server/handle_server_create_account.go +++ b/server/handle_server_create_account.go @@ -25,7 +25,7 @@ type ComAtprotoServerCreateAccountRequest struct { Handle string `json:"handle" validate:"required,atproto-handle"` Did *string `json:"did" validate:"atproto-did"` Password string `json:"password" validate:"required"` - InviteCode string `json:"inviteCode" validate:"required"` + InviteCode string `json:"inviteCode" validate:"omitempty"` } type ComAtprotoServerCreateAccountResponse struct { @@ -104,16 +104,22 @@ func (s *Server) handleCreateAccount(e echo.Context) error { } var ic models.InviteCode - if err := s.db.Raw("SELECT * FROM invite_codes WHERE code = ?", nil, request.InviteCode).Scan(&ic).Error; err != nil { - if err == gorm.ErrRecordNotFound { + if s.config.RequireInvite { + if strings.TrimSpace(request.InviteCode) == "" { return helpers.InputError(e, to.StringPtr("InvalidInviteCode")) } - s.logger.Error("error getting invite code from db", "error", err) - return helpers.ServerError(e, nil) - } - if ic.RemainingUseCount < 1 { - return helpers.InputError(e, to.StringPtr("InvalidInviteCode")) + if err := s.db.Raw("SELECT * FROM invite_codes WHERE code = ?", nil, request.InviteCode).Scan(&ic).Error; err != nil { + if err == gorm.ErrRecordNotFound { + return helpers.InputError(e, to.StringPtr("InvalidInviteCode")) + } + s.logger.Error("error getting invite code from db", "error", err) + return helpers.ServerError(e, nil) + } + + if ic.RemainingUseCount < 1 { + return helpers.InputError(e, to.StringPtr("InvalidInviteCode")) + } } // see if the email is already taken @@ -234,9 +240,11 @@ func (s *Server) handleCreateAccount(e echo.Context) error { }) } - if err := s.db.Raw("UPDATE invite_codes SET remaining_use_count = remaining_use_count - 1 WHERE code = ?", nil, request.InviteCode).Scan(&ic).Error; err != nil { - s.logger.Error("error decrementing use count", "error", err) - return helpers.ServerError(e, nil) + if s.config.RequireInvite { + if err := s.db.Raw("UPDATE invite_codes SET remaining_use_count = remaining_use_count - 1 WHERE code = ?", nil, request.InviteCode).Scan(&ic).Error; err != nil { + s.logger.Error("error decrementing use count", "error", err) + return helpers.ServerError(e, nil) + } } sess, err := s.createSession(&urepo) diff --git a/server/handle_server_describe_server.go b/server/handle_server_describe_server.go index 72e5628..c7a224a 100644 --- a/server/handle_server_describe_server.go +++ b/server/handle_server_describe_server.go @@ -22,7 +22,7 @@ type ComAtprotoServerDescribeServerResponse struct { func (s *Server) handleDescribeServer(e echo.Context) error { return e.JSON(200, ComAtprotoServerDescribeServerResponse{ - InviteCodeRequired: true, + InviteCodeRequired: s.config.RequireInvite, PhoneVerificationRequired: false, AvailableUserDomains: []string{"." + s.config.Hostname}, // TODO: more Links: ComAtprotoServerDescribeServerResponseLinks{ diff --git a/server/handle_well_known.go b/server/handle_well_known.go index 3419df1..565bbeb 100644 --- a/server/handle_well_known.go +++ b/server/handle_well_known.go @@ -2,9 +2,12 @@ package server import ( "fmt" + "strings" "github.com/Azure/go-autorest/autorest/to" + "github.com/haileyok/cocoon/internal/helpers" "github.com/labstack/echo/v4" + "gorm.io/gorm" ) var ( @@ -63,6 +66,36 @@ func (s *Server) handleWellKnown(e echo.Context) error { }) } +func (s *Server) handleAtprotoDid(e echo.Context) error { + host := e.Request().Host + if host == "" { + return helpers.InputError(e, to.StringPtr("Invalid handle.")) + } + + host = strings.Split(host, ":")[0] + host = strings.ToLower(strings.TrimSpace(host)) + + if host == s.config.Hostname { + return e.String(200, s.config.Did) + } + + suffix := "." + s.config.Hostname + if !strings.HasSuffix(host, suffix) { + return e.NoContent(404) + } + + actor, err := s.getActorByHandle(host) + if err != nil { + if err == gorm.ErrRecordNotFound { + return e.NoContent(404) + } + s.logger.Error("error looking up actor by handle", "error", err) + return helpers.ServerError(e, nil) + } + + return e.String(200, actor.Did) +} + func (s *Server) handleOauthProtectedResource(e echo.Context) error { return e.JSON(200, map[string]any{ "resource": "https://" + s.config.Hostname, diff --git a/server/server.go b/server/server.go index e849b6b..49340f9 100644 --- a/server/server.go +++ b/server/server.go @@ -102,6 +102,7 @@ type Args struct { ContactEmail string Relays []string AdminPassword string + RequireInvite bool SmtpUser string SmtpPass string @@ -126,6 +127,7 @@ type config struct { EnforcePeering bool Relays []string AdminPassword string + RequireInvite bool SmtpEmail string SmtpName string BlockstoreVariant BlockstoreVariant @@ -379,6 +381,7 @@ func New(args *Args) (*Server, error) { EnforcePeering: false, Relays: args.Relays, AdminPassword: args.AdminPassword, + RequireInvite: args.RequireInvite, SmtpName: args.SmtpName, SmtpEmail: args.SmtpEmail, BlockstoreVariant: args.BlockstoreVariant, @@ -442,6 +445,7 @@ func (s *Server) addRoutes() { s.echo.GET("/", s.handleRoot) s.echo.GET("/xrpc/_health", s.handleHealth) s.echo.GET("/.well-known/did.json", s.handleWellKnown) + s.echo.GET("/.well-known/atproto-did", s.handleAtprotoDid) s.echo.GET("/.well-known/oauth-protected-resource", s.handleOauthProtectedResource) s.echo.GET("/.well-known/oauth-authorization-server", s.handleOauthAuthorizationServer) s.echo.GET("/robots.txt", s.handleRobots) @@ -466,6 +470,9 @@ func (s *Server) addRoutes() { s.echo.GET("/xrpc/com.atproto.sync.listBlobs", s.handleSyncListBlobs) s.echo.GET("/xrpc/com.atproto.sync.getBlob", s.handleSyncGetBlob) + // labels + s.echo.GET("/xrpc/com.atproto.label.queryLabels", s.handleLabelQueryLabels) + // account s.echo.GET("/account", s.handleAccount) s.echo.POST("/account/revoke", s.handleAccountRevoke)