diff --git a/internal/api/handlers/actor/errors.go b/internal/api/handlers/actor/errors.go index ebbc8eb..77b38a0 100644 --- a/internal/api/handlers/actor/errors.go +++ b/internal/api/handlers/actor/errors.go @@ -1,71 +1,46 @@ package actor import ( - "encoding/json" - "errors" "fmt" - "log" "net/http" + "Coves/internal/api/xrpc" "Coves/internal/core/posts" ) -// ErrorResponse represents an XRPC error response -type ErrorResponse struct { - Error string `json:"error"` - Message string `json:"message"` -} +// ErrorResponse represents an XRPC error response. +type ErrorResponse = xrpc.Error + +// errorMapper maps actor service errors to XRPC responses. +// +// Every way an actor can be missing answers ActorNotFound: this package reads +// posts for a specific actor, so a missing post here means the actor is what we +// could not find. +var errorMapper = xrpc.NewMapper("actor", + xrpc.As[*actorNotFoundError](http.StatusNotFound, "ActorNotFound", + func(*actorNotFoundError) string { return "Actor not found" }), + xrpc.Sentinel(posts.ErrNotFound, http.StatusNotFound, + "ActorNotFound", "Actor not found"), + xrpc.Sentinel(posts.ErrActorNotFound, http.StatusNotFound, + "ActorNotFound", "Actor not found"), + xrpc.Sentinel(posts.ErrCommunityNotFound, http.StatusNotFound, + "CommunityNotFound", "Community not found"), + xrpc.Sentinel(posts.ErrInvalidCursor, http.StatusBadRequest, + "InvalidCursor", "Invalid pagination cursor"), + // Message comes off the typed error rather than the chain, so context added + // by a caller cannot leak into the response. + xrpc.As[*posts.ValidationError](http.StatusBadRequest, "InvalidRequest", + func(e *posts.ValidationError) string { return e.Message }), +) -// writeError writes a JSON error response +// writeError writes a JSON error response. func writeError(w http.ResponseWriter, statusCode int, errorType, message string) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(statusCode) - if err := json.NewEncoder(w).Encode(ErrorResponse{ - Error: errorType, - Message: message, - }); err != nil { - // Log encoding errors but can't send error response (headers already sent) - log.Printf("ERROR: Failed to encode error response: %v", err) - } + xrpc.WriteError(w, statusCode, errorType, message) } -// handleServiceError maps service errors to HTTP responses +// handleServiceError maps service errors to HTTP responses. func handleServiceError(w http.ResponseWriter, err error) { - // Check for handler-level errors first - var actorNotFound *actorNotFoundError - if errors.As(err, &actorNotFound) { - writeError(w, http.StatusNotFound, "ActorNotFound", "Actor not found") - return - } - - // Check for service-level errors - switch { - case errors.Is(err, posts.ErrNotFound): - writeError(w, http.StatusNotFound, "ActorNotFound", "Actor not found") - - case errors.Is(err, posts.ErrActorNotFound): - writeError(w, http.StatusNotFound, "ActorNotFound", "Actor not found") - - case errors.Is(err, posts.ErrCommunityNotFound): - writeError(w, http.StatusNotFound, "CommunityNotFound", "Community not found") - - case errors.Is(err, posts.ErrInvalidCursor): - writeError(w, http.StatusBadRequest, "InvalidCursor", "Invalid pagination cursor") - - case posts.IsValidationError(err): - // Extract message from ValidationError for cleaner response - var valErr *posts.ValidationError - if errors.As(err, &valErr) { - writeError(w, http.StatusBadRequest, "InvalidRequest", valErr.Message) - } else { - writeError(w, http.StatusBadRequest, "InvalidRequest", err.Error()) - } - - default: - // Internal server error - don't leak details - log.Printf("ERROR: Actor posts service error: %v", err) - writeError(w, http.StatusInternalServerError, "InternalServerError", "An internal error occurred") - } + errorMapper.Write(w, err) } // actorNotFoundError represents an actor not found error diff --git a/internal/api/handlers/adminreport/errors.go b/internal/api/handlers/adminreport/errors.go index ef6971d..e236d6a 100644 --- a/internal/api/handlers/adminreport/errors.go +++ b/internal/api/handlers/adminreport/errors.go @@ -1,53 +1,45 @@ package adminreport import ( + "net/http" + "Coves/internal/api/xrpc" "Coves/internal/core/adminreports" - "errors" - "log" - "net/http" ) -// writeError writes a JSON error response with the given status code +// errorMapper maps admin report service errors to XRPC responses. +// +// Each validation sentinel gets a static message: the enumerated values in +// these are the useful part, and spelling them out beats echoing the error. +var errorMapper = xrpc.NewMapper("adminreport", + xrpc.Sentinel(adminreports.ErrInvalidReason, http.StatusBadRequest, "InvalidReason", + "Invalid report reason. Must be one of: csam, doxing, harassment, spam, illegal, other"), + xrpc.Sentinel(adminreports.ErrInvalidStatus, http.StatusBadRequest, "InvalidStatus", + "Invalid report status. Must be one of: open, reviewing, resolved, dismissed"), + xrpc.Sentinel(adminreports.ErrInvalidTarget, http.StatusBadRequest, "InvalidTarget", + "Invalid target URI. Must be a valid AT Protocol URI starting with at://"), + xrpc.Sentinel(adminreports.ErrExplanationTooLong, http.StatusBadRequest, "ExplanationTooLong", + "Explanation exceeds maximum length of 1000 characters"), + xrpc.Sentinel(adminreports.ErrReporterRequired, http.StatusBadRequest, "ReporterRequired", + "Reporter DID is required"), + xrpc.Sentinel(adminreports.ErrInvalidTargetType, http.StatusBadRequest, "InvalidTargetType", + "Invalid target type. Must be one of: post, comment"), + + xrpc.Match(adminreports.IsNotFound, http.StatusNotFound, + "NotFound", "Report not found"), + + // Catch-all so a validation sentinel added to the domain without a rule + // above still answers 400 rather than 500. + xrpc.Match(adminreports.IsValidationError, http.StatusBadRequest, + "InvalidRequest", "The request contains invalid data"), +) + +// writeError writes a JSON error response with the given status code. func writeError(w http.ResponseWriter, statusCode int, errorType, message string) { xrpc.WriteError(w, statusCode, errorType, message) } -// handleServiceError maps service-layer errors to HTTP responses +// handleServiceError maps service-layer errors to HTTP responses. func handleServiceError(w http.ResponseWriter, err error) { - switch { - case adminreports.IsValidationError(err): - switch { - case errors.Is(err, adminreports.ErrInvalidReason): - writeError(w, http.StatusBadRequest, "InvalidReason", - "Invalid report reason. Must be one of: csam, doxing, harassment, spam, illegal, other") - case errors.Is(err, adminreports.ErrInvalidStatus): - writeError(w, http.StatusBadRequest, "InvalidStatus", - "Invalid report status. Must be one of: open, reviewing, resolved, dismissed") - case errors.Is(err, adminreports.ErrInvalidTarget): - writeError(w, http.StatusBadRequest, "InvalidTarget", - "Invalid target URI. Must be a valid AT Protocol URI starting with at://") - case errors.Is(err, adminreports.ErrExplanationTooLong): - writeError(w, http.StatusBadRequest, "ExplanationTooLong", - "Explanation exceeds maximum length of 1000 characters") - case errors.Is(err, adminreports.ErrReporterRequired): - writeError(w, http.StatusBadRequest, "ReporterRequired", - "Reporter DID is required") - case errors.Is(err, adminreports.ErrInvalidTargetType): - writeError(w, http.StatusBadRequest, "InvalidTargetType", - "Invalid target type. Must be one of: post, comment") - default: - log.Printf("Unhandled validation error in admin report handler: %v", err) - writeError(w, http.StatusBadRequest, "InvalidRequest", - "The request contains invalid data") - } - - case adminreports.IsNotFound(err): - writeError(w, http.StatusNotFound, "NotFound", "Report not found") - - default: - log.Printf("Unexpected error in admin report handler: %v", err) - writeError(w, http.StatusInternalServerError, "InternalServerError", - "An internal error occurred") - } + errorMapper.Write(w, err) } diff --git a/internal/api/handlers/aggregator/errors.go b/internal/api/handlers/aggregator/errors.go index 579fa64..cc016be 100644 --- a/internal/api/handlers/aggregator/errors.go +++ b/internal/api/handlers/aggregator/errors.go @@ -1,19 +1,15 @@ package aggregator import ( - "Coves/internal/core/aggregators" - "Coves/internal/core/communities" "bytes" "encoding/json" "log" "net/http" -) -// ErrorResponse represents an XRPC error response -type ErrorResponse struct { - Error string `json:"error"` - Message string `json:"message"` -} + "Coves/internal/api/xrpc" + "Coves/internal/core/aggregators" + "Coves/internal/core/communities" +) // writeJSONResponse buffers the JSON encoding before sending headers. // This ensures that encoding failures don't result in partial responses @@ -40,44 +36,31 @@ func writeJSONResponse(w http.ResponseWriter, statusCode int, data interface{}) return true } -// writeError writes a JSON error response with proper buffering +// errorMapper maps aggregator service errors to XRPC responses. +// +// Handlers here call into communities as well, to resolve a community +// identifier, so those errors are checked first and answer with the more +// specific CommunityNotFound. +var errorMapper = xrpc.NewMapper("aggregator", + xrpc.MatchDetail(communities.IsNotFound, http.StatusNotFound, "CommunityNotFound"), + xrpc.MatchDetail(communities.IsValidationError, http.StatusBadRequest, "InvalidRequest"), + + xrpc.MatchDetail(aggregators.IsNotFound, http.StatusNotFound, "NotFound"), + xrpc.MatchDetail(aggregators.IsValidationError, http.StatusBadRequest, "InvalidRequest"), + xrpc.MatchDetail(aggregators.IsUnauthorized, http.StatusForbidden, "Forbidden"), + xrpc.MatchDetail(aggregators.IsConflict, http.StatusConflict, "Conflict"), + xrpc.MatchDetail(aggregators.IsRateLimited, http.StatusTooManyRequests, "RateLimitExceeded"), + xrpc.Match(aggregators.IsNotImplemented, http.StatusNotImplemented, + "NotImplemented", "This feature is not yet available (Phase 2)"), +) + +// writeError writes a JSON error response with proper buffering. func writeError(w http.ResponseWriter, statusCode int, errorType, message string) { - writeJSONResponse(w, statusCode, ErrorResponse{ - Error: errorType, - Message: message, - }) + xrpc.WriteError(w, statusCode, errorType, message) } -// handleServiceError maps service errors to HTTP responses -// Handles errors from both aggregators and communities packages +// handleServiceError maps service errors to HTTP responses. +// Handles errors from both aggregators and communities packages. func handleServiceError(w http.ResponseWriter, err error) { - if err == nil { - return - } - - // Map domain errors to HTTP status codes - // Check community errors first (for ResolveCommunityIdentifier calls) - switch { - case communities.IsNotFound(err): - writeError(w, http.StatusNotFound, "CommunityNotFound", err.Error()) - case communities.IsValidationError(err): - writeError(w, http.StatusBadRequest, "InvalidRequest", err.Error()) - case aggregators.IsNotFound(err): - writeError(w, http.StatusNotFound, "NotFound", err.Error()) - case aggregators.IsValidationError(err): - writeError(w, http.StatusBadRequest, "InvalidRequest", err.Error()) - case aggregators.IsUnauthorized(err): - writeError(w, http.StatusForbidden, "Forbidden", err.Error()) - case aggregators.IsConflict(err): - writeError(w, http.StatusConflict, "Conflict", err.Error()) - case aggregators.IsRateLimited(err): - writeError(w, http.StatusTooManyRequests, "RateLimitExceeded", err.Error()) - case aggregators.IsNotImplemented(err): - writeError(w, http.StatusNotImplemented, "NotImplemented", "This feature is not yet available (Phase 2)") - default: - // Internal errors - don't leak details - log.Printf("ERROR: Aggregator service error: %v", err) - writeError(w, http.StatusInternalServerError, "InternalServerError", - "An internal error occurred") - } + errorMapper.Write(w, err) } diff --git a/internal/api/handlers/comments/errors.go b/internal/api/handlers/comments/errors.go index 117db47..0088cf2 100644 --- a/internal/api/handlers/comments/errors.go +++ b/internal/api/handlers/comments/errors.go @@ -1,83 +1,69 @@ package comments import ( - "Coves/internal/core/comments" - "encoding/json" - "errors" - "log" "net/http" -) -// errorResponse represents a standardized JSON error response -type errorResponse struct { - Error string `json:"error"` - Message string `json:"message"` -} + "Coves/internal/api/xrpc" + "Coves/internal/core/comments" +) -// writeError writes a JSON error response with the given status code -func writeError(w http.ResponseWriter, statusCode int, errorType, message string) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(statusCode) - if err := json.NewEncoder(w).Encode(errorResponse{ - Error: errorType, - Message: message, - }); err != nil { - log.Printf("Failed to encode error response: %v", err) - } -} +// errorMapper maps comment service errors to XRPC responses. +// +// comments signals validation and not-found with sentinels rather than the +// shared typed errors, so unlike most packages it spells out both groups here. +// Everything else — dead sessions, other PDS failures, request lifecycle — +// comes from xrpc's shared rules. +var errorMapper = xrpc.NewMapper("comments", + // Not found. + xrpc.Sentinel(comments.ErrCommentNotFound, http.StatusNotFound, + "CommentNotFound", "Comment not found"), + xrpc.Sentinel(comments.ErrParentNotFound, http.StatusNotFound, + "ParentNotFound", "Parent post or comment not found"), + xrpc.Sentinel(comments.ErrRootNotFound, http.StatusNotFound, + "RootNotFound", "Root post not found"), -// handleServiceError maps service-layer errors to HTTP responses -// This follows the error handling pattern from other handlers (post, community) -func handleServiceError(w http.ResponseWriter, err error) { - switch { - case comments.IsNotFound(err): - // Map specific not found errors to appropriate messages - switch { - case errors.Is(err, comments.ErrCommentNotFound): - writeError(w, http.StatusNotFound, "CommentNotFound", "Comment not found") - case errors.Is(err, comments.ErrParentNotFound): - writeError(w, http.StatusNotFound, "ParentNotFound", "Parent post or comment not found") - case errors.Is(err, comments.ErrRootNotFound): - writeError(w, http.StatusNotFound, "RootNotFound", "Root post not found") - default: - writeError(w, http.StatusNotFound, "NotFound", err.Error()) - } + // Validation. + xrpc.Sentinel(comments.ErrInvalidReply, http.StatusBadRequest, + "InvalidReply", "The reply reference is invalid or malformed"), + xrpc.Sentinel(comments.ErrContentTooLong, http.StatusBadRequest, + "ContentTooLong", "Comment content exceeds 10000 graphemes"), + xrpc.Sentinel(comments.ErrContentEmpty, http.StatusBadRequest, + "ContentEmpty", "Comment content is required"), + // The error names which facet and which field, which is client-actionable + // and carries no internal state. + xrpc.SentinelDetail(comments.ErrInvalidFacets, http.StatusBadRequest, + "InvalidFacets"), + // Fixed message: the repository layer wraps this one with detail the client + // has no use for. + xrpc.Sentinel(comments.ErrInvalidCursor, http.StatusBadRequest, + "InvalidRequest", "Invalid or mismatched pagination cursor"), - case comments.IsValidationError(err): - // Map specific validation errors to appropriate messages - switch { - case errors.Is(err, comments.ErrInvalidReply): - writeError(w, http.StatusBadRequest, "InvalidReply", "The reply reference is invalid or malformed") - case errors.Is(err, comments.ErrContentTooLong): - writeError(w, http.StatusBadRequest, "ContentTooLong", "Comment content exceeds 10000 graphemes") - case errors.Is(err, comments.ErrContentEmpty): - writeError(w, http.StatusBadRequest, "ContentEmpty", "Comment content is required") - case errors.Is(err, comments.ErrInvalidFacets): - // err carries the structural detail (which facet, which field); it is - // client-actionable and contains no internal state - writeError(w, http.StatusBadRequest, "InvalidFacets", err.Error()) - case errors.Is(err, comments.ErrInvalidCursor): - // Fixed message avoids leaking internal wrapping detail from the repository layer - writeError(w, http.StatusBadRequest, "InvalidRequest", "Invalid or mismatched pagination cursor") - default: - writeError(w, http.StatusBadRequest, "InvalidRequest", err.Error()) - } + // A concurrent edit lost the optimistic-locking race on PutRecord. This + // used to be absent, on the theory that the PDS never surfaced a conflict — + // it does, and the omission turned every lost race into a 500. + xrpc.Sentinel(comments.ErrConcurrentModification, http.StatusConflict, + "ConcurrentModification", "The comment was modified by another request. Fetch it again and retry."), - case errors.Is(err, comments.ErrNotAuthorized): - writeError(w, http.StatusForbidden, "NotAuthorized", "User is not authorized to perform this action") + // Authorization. + xrpc.Sentinel(comments.ErrNotAuthorized, http.StatusForbidden, + "NotAuthorized", "User is not authorized to perform this action"), + xrpc.Sentinel(comments.ErrBanned, http.StatusForbidden, + "Banned", "User is banned from this community"), - case errors.Is(err, comments.ErrBanned): - writeError(w, http.StatusForbidden, "Banned", "User is banned from this community") + // Catch-alls for sentinels added to the domain without a rule here: a new + // not-found still answers 404 rather than 500. + xrpc.Match(comments.IsNotFound, http.StatusNotFound, + "NotFound", "The requested resource was not found"), + xrpc.Match(comments.IsValidationError, http.StatusBadRequest, + "InvalidRequest", "The request contains invalid data"), +) - // NOTE: IsConflict case removed - the PDS handles duplicate detection via CreateRecord, - // so ErrCommentAlreadyExists is never returned from the service layer. If the PDS rejects - // a duplicate record, it returns an auth/validation error which is handled by other cases. - // Keeping this code would be dead code that never executes. +// writeError writes a JSON error response with the given status code. +func writeError(w http.ResponseWriter, statusCode int, errorType, message string) { + xrpc.WriteError(w, statusCode, errorType, message) +} - default: - // Don't leak internal error details to clients - log.Printf("Unexpected error in comments handler: %v", err) - writeError(w, http.StatusInternalServerError, "InternalServerError", - "An internal error occurred") - } +// handleServiceError maps service-layer errors to HTTP responses. +func handleServiceError(w http.ResponseWriter, err error) { + errorMapper.Write(w, err) } diff --git a/internal/api/handlers/comments/errors_test.go b/internal/api/handlers/comments/errors_test.go new file mode 100644 index 0000000..8252b24 --- /dev/null +++ b/internal/api/handlers/comments/errors_test.go @@ -0,0 +1,113 @@ +package comments + +import ( + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "testing" + + "Coves/internal/api/xrpc" + "Coves/internal/atproto/pds" + "Coves/internal/core/comments" +) + +// Commenting write-forwards to the user's repo, so an expired session used to +// answer 500 here with no signal for the client to re-authenticate. +func TestExpiredSessionAnswers401(t *testing.T) { + tests := []struct { + name string + err error + wantStatus int + wantCode string + }{ + { + name: "pds rejected the token", + err: fmt.Errorf("%w: %w", comments.ErrNotAuthorized, + fmt.Errorf("CreateRecord: %w: expired", pds.ErrUnauthorized)), + wantStatus: http.StatusUnauthorized, + wantCode: "AuthRequired", + }, + { + name: "session could not be resumed", + err: fmt.Errorf("failed to create PDS client: %w", + fmt.Errorf("resume: %w", pds.ErrSessionExpired)), + wantStatus: http.StatusUnauthorized, + wantCode: "AuthRequired", + }, + { + name: "missing scope stays 403", + err: fmt.Errorf("%w: %w", comments.ErrNotAuthorized, + fmt.Errorf("CreateRecord: %w", pds.ErrForbidden)), + wantStatus: http.StatusForbidden, + wantCode: "NotAuthorized", + }, + { + name: "appview refusal stays 403", + err: comments.ErrNotAuthorized, + wantStatus: http.StatusForbidden, + wantCode: "NotAuthorized", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + rec := httptest.NewRecorder() + handleServiceError(rec, tt.err) + assertXRPCError(t, rec, tt.wantStatus, tt.wantCode) + }) + } +} + +// A concurrent edit loses the optimistic-locking race on PutRecord and the +// service reports ErrConcurrentModification. The old switch dropped that case +// as "dead code" it was not, so every lost race answered 500. +func TestConcurrentModificationIs409(t *testing.T) { + rec := httptest.NewRecorder() + handleServiceError(rec, fmt.Errorf("%w: %w", comments.ErrConcurrentModification, + fmt.Errorf("PutRecord: %w", pds.ErrConflict))) + assertXRPCError(t, rec, http.StatusConflict, "ConcurrentModification") +} + +// The comment error codes are a client contract; this pins them. +func TestCommentErrorCodesUnchanged(t *testing.T) { + tests := []struct { + err error + wantStatus int + wantCode string + }{ + {comments.ErrCommentNotFound, http.StatusNotFound, "CommentNotFound"}, + {comments.ErrParentNotFound, http.StatusNotFound, "ParentNotFound"}, + {comments.ErrRootNotFound, http.StatusNotFound, "RootNotFound"}, + {comments.ErrInvalidReply, http.StatusBadRequest, "InvalidReply"}, + {comments.ErrContentTooLong, http.StatusBadRequest, "ContentTooLong"}, + {comments.ErrContentEmpty, http.StatusBadRequest, "ContentEmpty"}, + {comments.ErrInvalidFacets, http.StatusBadRequest, "InvalidFacets"}, + {comments.ErrInvalidCursor, http.StatusBadRequest, "InvalidRequest"}, + {comments.ErrBanned, http.StatusForbidden, "Banned"}, + } + + for _, tt := range tests { + t.Run(tt.wantCode+"/"+tt.err.Error(), func(t *testing.T) { + rec := httptest.NewRecorder() + handleServiceError(rec, fmt.Errorf("service: %w", tt.err)) + assertXRPCError(t, rec, tt.wantStatus, tt.wantCode) + }) + } +} + +func assertXRPCError(t *testing.T, rec *httptest.ResponseRecorder, wantStatus int, wantCode string) xrpc.Error { + t.Helper() + + if rec.Code != wantStatus { + t.Errorf("status = %d, want %d", rec.Code, wantStatus) + } + var body xrpc.Error + if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { + t.Fatalf("body is not valid JSON: %v", err) + } + if body.Error != wantCode { + t.Errorf("code = %q, want %q", body.Error, wantCode) + } + return body +} diff --git a/internal/api/handlers/community/errors.go b/internal/api/handlers/community/errors.go index 7b9c72e..e6d5365 100644 --- a/internal/api/handlers/community/errors.go +++ b/internal/api/handlers/community/errors.go @@ -1,66 +1,42 @@ package community import ( - "Coves/internal/atproto/pds" - "Coves/internal/core/communities" - "encoding/json" - "errors" - "log" "net/http" + + "Coves/internal/api/xrpc" + "Coves/internal/core/communities" ) -// XRPCError represents an XRPC error response -type XRPCError struct { - Error string `json:"error"` - Message string `json:"message"` -} +// errorMapper maps community service errors to XRPC responses. +// +// The PDS rules this package used to spell out by hand now come from xrpc's +// shared rules, along with dead sessions and request-lifecycle errors. That +// also closed two gaps: a rate-limited or oversized write to the PDS answered +// 500 here, because those two sentinels were the ones the hand-written switch +// happened to omit. +var errorMapper = xrpc.NewMapper("community", + // Ahead of the generic conflict rule, which matches this sentinel too. + xrpc.Sentinel(communities.ErrHandleTaken, http.StatusConflict, + "NameTaken", "Community handle is already taken"), + + xrpc.Sentinel(communities.ErrUnauthorized, http.StatusForbidden, + "Forbidden", "You do not have permission to perform this action"), + xrpc.Sentinel(communities.ErrMemberBanned, http.StatusForbidden, + "Blocked", "You are blocked from this community"), + + // The domain's own predicates. Their sentinel text is written for the + // client, so it doubles as the message. + xrpc.MatchDetail(communities.IsNotFound, http.StatusNotFound, "NotFound"), + xrpc.MatchDetail(communities.IsConflict, http.StatusConflict, "AlreadyExists"), + xrpc.MatchDetail(communities.IsValidationError, http.StatusBadRequest, "InvalidRequest"), +) -// writeError writes an XRPC error response -func writeError(w http.ResponseWriter, status int, error, message string) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(status) - if err := json.NewEncoder(w).Encode(XRPCError{ - Error: error, - Message: message, - }); err != nil { - log.Printf("Failed to encode error response: %v", err) - } +// writeError writes an XRPC error response. +func writeError(w http.ResponseWriter, status int, code, message string) { + xrpc.WriteError(w, status, code, message) } -// handleServiceError converts service errors to appropriate HTTP responses +// handleServiceError converts service errors to appropriate HTTP responses. func handleServiceError(w http.ResponseWriter, err error) { - switch { - case communities.IsNotFound(err): - writeError(w, http.StatusNotFound, "NotFound", err.Error()) - case communities.IsConflict(err): - if err == communities.ErrHandleTaken { - writeError(w, http.StatusConflict, "NameTaken", "Community handle is already taken") - } else { - writeError(w, http.StatusConflict, "AlreadyExists", err.Error()) - } - case communities.IsValidationError(err): - writeError(w, http.StatusBadRequest, "InvalidRequest", err.Error()) - case err == communities.ErrUnauthorized: - writeError(w, http.StatusForbidden, "Forbidden", "You do not have permission to perform this action") - case err == communities.ErrMemberBanned: - writeError(w, http.StatusForbidden, "Blocked", "You are blocked from this community") - // PDS-specific errors (from DPoP authentication or PDS API calls) - case errors.Is(err, pds.ErrBadRequest): - writeError(w, http.StatusBadRequest, "InvalidRequest", "Invalid request to PDS") - case errors.Is(err, pds.ErrNotFound): - writeError(w, http.StatusNotFound, "NotFound", "Record not found on PDS") - case errors.Is(err, pds.ErrConflict): - writeError(w, http.StatusConflict, "Conflict", "Record was modified by another operation") - case errors.Is(err, pds.ErrForbidden): - // 403 is a permissions problem (e.g. missing OAuth scope), not an - // expired session — it must not trigger a client sign-out. - writeError(w, http.StatusForbidden, "PermissionDenied", "Your session does not have permission for this action. Sign out and back in to grant it.") - case errors.Is(err, pds.ErrUnauthorized): - // PDS auth errors should prompt re-authentication - writeError(w, http.StatusUnauthorized, "AuthRequired", "Authentication required or session expired") - default: - // Internal server error - log the actual error for debugging - log.Printf("XRPC handler error: %v", err) - writeError(w, http.StatusInternalServerError, "InternalServerError", "An internal error occurred") - } + errorMapper.Write(w, err) } diff --git a/internal/api/handlers/community/errors_test.go b/internal/api/handlers/community/errors_test.go new file mode 100644 index 0000000..dfbfe7a --- /dev/null +++ b/internal/api/handlers/community/errors_test.go @@ -0,0 +1,151 @@ +package community + +import ( + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "Coves/internal/api/xrpc" + "Coves/internal/atproto/pds" + "Coves/internal/core/communities" +) + +// TestCommunityErrorCodes pins the wire contract, exercising every sentinel +// WRAPPED — which is how it actually arrives, since the service adds context on +// the way up. The mapper this replaced compared with ==, so a wrapped sentinel +// fell through; these rows are the behavior that changed. +// +// NOTE: ErrHandleTaken is a deliberate contract change. The old switch reached +// NameTaken only via `err == communities.ErrHandleTaken`, and the service wraps +// it ("failed to persist community with credentials: %w"), so the real create +// path always answered AlreadyExists. errors.Is now matches, so it answers +// NameTaken — the code the old switch plainly intended. Clients keying on +// AlreadyExists for a handle collision need updating. +func TestCommunityErrorCodes(t *testing.T) { + tests := []struct { + name string + err error + wantStatus int + wantCode string + }{ + { + name: "handle collision reaches NameTaken through a wrapper", + err: fmt.Errorf("failed to persist community: %w", communities.ErrHandleTaken), + wantStatus: http.StatusConflict, + wantCode: "NameTaken", + }, + { + name: "other conflicts stay AlreadyExists", + err: fmt.Errorf("service: %w", communities.ErrCommunityAlreadyExists), + wantStatus: http.StatusConflict, + wantCode: "AlreadyExists", + }, + { + name: "appview authorization refusal", + err: fmt.Errorf("service: %w", communities.ErrUnauthorized), + wantStatus: http.StatusForbidden, + wantCode: "Forbidden", + }, + { + name: "banned member", + err: fmt.Errorf("service: %w", communities.ErrMemberBanned), + wantStatus: http.StatusForbidden, + wantCode: "Blocked", + }, + { + name: "missing community", + err: fmt.Errorf("service: %w", communities.ErrCommunityNotFound), + wantStatus: http.StatusNotFound, + wantCode: "NotFound", + }, + { + name: "invalid input", + err: fmt.Errorf("service: %w", communities.ErrInvalidInput), + wantStatus: http.StatusBadRequest, + wantCode: "InvalidRequest", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + rec := httptest.NewRecorder() + handleServiceError(rec, tt.err) + assertXRPCError(t, rec, tt.wantStatus, tt.wantCode) + }) + } +} + +// Subscribe and block write to the user's own repo, so a dead session here must +// reach the client as 401 — this package's whole reason for having PDS rules. +func TestCommunityPDSErrors(t *testing.T) { + tests := []struct { + name string + err error + wantStatus int + wantCode string + }{ + { + name: "expired token behind the domain sentinel", + err: fmt.Errorf("%w: %w", communities.ErrUnauthorized, + fmt.Errorf("CreateRecord: %w", pds.ErrUnauthorized)), + wantStatus: http.StatusUnauthorized, + wantCode: "AuthRequired", + }, + { + name: "missing scope keeps the domain's 403", + err: fmt.Errorf("%w: %w", communities.ErrUnauthorized, + fmt.Errorf("CreateRecord: %w", pds.ErrForbidden)), + wantStatus: http.StatusForbidden, + wantCode: "Forbidden", + }, + { + // Previously 500: this package's hand-written switch omitted 429/413. + name: "pds rate limit", + err: fmt.Errorf("CreateRecord: %w", pds.ErrRateLimited), + wantStatus: http.StatusTooManyRequests, + wantCode: "RateLimitExceeded", + }, + { + name: "pds payload too large", + err: fmt.Errorf("CreateRecord: %w", pds.ErrPayloadTooLarge), + wantStatus: http.StatusRequestEntityTooLarge, + wantCode: "PayloadTooLarge", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + rec := httptest.NewRecorder() + handleServiceError(rec, tt.err) + assertXRPCError(t, rec, tt.wantStatus, tt.wantCode) + }) + } +} + +func TestUnmappedErrorIsGeneric500(t *testing.T) { + rec := httptest.NewRecorder() + handleServiceError(rec, fmt.Errorf("pq: connection refused to 10.0.0.4:5432")) + + assertXRPCError(t, rec, http.StatusInternalServerError, "InternalServerError") + if body := rec.Body.String(); strings.Contains(body, "10.0.0.4") { + t.Errorf("internal detail leaked to the client: %s", body) + } +} + +func assertXRPCError(t *testing.T, rec *httptest.ResponseRecorder, wantStatus int, wantCode string) { + t.Helper() + + if rec.Code != wantStatus { + t.Errorf("status = %d, want %d", rec.Code, wantStatus) + } + var body xrpc.Error + if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { + t.Fatalf("body is not valid JSON: %v", err) + } + if body.Error != wantCode { + t.Errorf("code = %q, want %q", body.Error, wantCode) + } +} diff --git a/internal/api/handlers/communityFeed/errors.go b/internal/api/handlers/communityFeed/errors.go index d51bd4e..263082a 100644 --- a/internal/api/handlers/communityFeed/errors.go +++ b/internal/api/handlers/communityFeed/errors.go @@ -1,46 +1,27 @@ package communityFeed import ( - "Coves/internal/core/communityFeeds" - "encoding/json" - "errors" - "log" "net/http" + + "Coves/internal/api/xrpc" + "Coves/internal/core/communityFeeds" ) -// ErrorResponse represents an XRPC error response -type ErrorResponse struct { - Error string `json:"error"` - Message string `json:"message"` -} +// errorMapper maps community feed service errors to XRPC responses. +var errorMapper = xrpc.NewMapper("communityFeed", + xrpc.Sentinel(communityFeeds.ErrCommunityNotFound, http.StatusNotFound, + "CommunityNotFound", "Community not found"), + xrpc.Sentinel(communityFeeds.ErrInvalidCursor, http.StatusBadRequest, + "InvalidCursor", "Invalid pagination cursor"), + xrpc.MatchDetail(communityFeeds.IsValidationError, http.StatusBadRequest, "InvalidRequest"), +) -// writeError writes a JSON error response +// writeError writes a JSON error response. func writeError(w http.ResponseWriter, statusCode int, errorType, message string) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(statusCode) - if err := json.NewEncoder(w).Encode(ErrorResponse{ - Error: errorType, - Message: message, - }); err != nil { - // Log encoding errors but can't send error response (headers already sent) - log.Printf("ERROR: Failed to encode error response: %v", err) - } + xrpc.WriteError(w, statusCode, errorType, message) } -// handleServiceError maps service errors to HTTP responses +// handleServiceError maps service errors to HTTP responses. func handleServiceError(w http.ResponseWriter, err error) { - switch { - case errors.Is(err, communityFeeds.ErrCommunityNotFound): - writeError(w, http.StatusNotFound, "CommunityNotFound", "Community not found") - - case errors.Is(err, communityFeeds.ErrInvalidCursor): - writeError(w, http.StatusBadRequest, "InvalidCursor", "Invalid pagination cursor") - - case communityFeeds.IsValidationError(err): - writeError(w, http.StatusBadRequest, "InvalidRequest", err.Error()) - - default: - // Internal server error - don't leak details - writeError(w, http.StatusInternalServerError, "InternalServerError", "An internal error occurred") - } + errorMapper.Write(w, err) } diff --git a/internal/api/handlers/communitysuggestion/errors.go b/internal/api/handlers/communitysuggestion/errors.go index 830e79c..1cf58f4 100644 --- a/internal/api/handlers/communitysuggestion/errors.go +++ b/internal/api/handlers/communitysuggestion/errors.go @@ -1,55 +1,55 @@ package communitysuggestion import ( + "net/http" + "Coves/internal/api/xrpc" "Coves/internal/core/communitysuggestions" - "errors" - "log" - "net/http" ) -// writeError writes an XRPC error response -func writeError(w http.ResponseWriter, status int, error, message string) { - xrpc.WriteError(w, status, error, message) +// errorMapper maps community suggestion service errors to XRPC responses. +// +// Every sentinel gets a static, user-facing message rather than echoing the +// error, so nothing internal reaches the client. +var errorMapper = xrpc.NewMapper("communitysuggestion", + xrpc.Sentinel(communitysuggestions.ErrTitleRequired, http.StatusBadRequest, + "InvalidRequest", "Suggestion title is required"), + xrpc.Sentinel(communitysuggestions.ErrTitleTooLong, http.StatusBadRequest, + "InvalidRequest", "Suggestion title exceeds maximum length"), + xrpc.Sentinel(communitysuggestions.ErrDescriptionRequired, http.StatusBadRequest, + "InvalidRequest", "Suggestion description is required"), + xrpc.Sentinel(communitysuggestions.ErrDescriptionTooLong, http.StatusBadRequest, + "InvalidRequest", "Suggestion description exceeds maximum length"), + xrpc.Sentinel(communitysuggestions.ErrInvalidStatus, http.StatusBadRequest, + "InvalidStatus", "Invalid status value. Must be one of: open, under_review, approved, declined"), + xrpc.Sentinel(communitysuggestions.ErrInvalidVoteValue, http.StatusBadRequest, + "InvalidVoteValue", "Invalid vote value. Must be 1 or -1"), + xrpc.Sentinel(communitysuggestions.ErrInvalidSuggestionID, http.StatusBadRequest, + "InvalidRequest", "Invalid suggestion ID"), + xrpc.Sentinel(communitysuggestions.ErrVoterRequired, http.StatusBadRequest, + "InvalidRequest", "Voter identification is required"), + xrpc.Sentinel(communitysuggestions.ErrSubmitterRequired, http.StatusBadRequest, + "InvalidRequest", "Submitter identification is required"), + + xrpc.Match(communitysuggestions.IsNotFound, http.StatusNotFound, + "NotFound", "The requested resource was not found"), + xrpc.Match(communitysuggestions.IsRateLimitError, http.StatusTooManyRequests, + "RateLimitExceeded", "Too many suggestions. Please try again later"), + xrpc.Match(communitysuggestions.IsAuthorizationError, http.StatusForbidden, + "Forbidden", "You are not authorized to perform this action"), + + // Catch-all so a validation sentinel added to the domain without a rule + // above still answers 400 rather than 500. + xrpc.Match(communitysuggestions.IsValidationError, http.StatusBadRequest, + "InvalidRequest", "The request contains invalid data"), +) + +// writeError writes an XRPC error response. +func writeError(w http.ResponseWriter, status int, code, message string) { + xrpc.WriteError(w, status, code, message) } -// handleServiceError converts service errors to appropriate HTTP responses -// Each sentinel error is mapped to a static, user-facing message to prevent -// leaking internal error details to clients. +// handleServiceError converts service errors to appropriate HTTP responses. func handleServiceError(w http.ResponseWriter, err error) { - switch { - case communitysuggestions.IsValidationError(err): - switch { - case errors.Is(err, communitysuggestions.ErrTitleRequired): - writeError(w, http.StatusBadRequest, "InvalidRequest", "Suggestion title is required") - case errors.Is(err, communitysuggestions.ErrTitleTooLong): - writeError(w, http.StatusBadRequest, "InvalidRequest", "Suggestion title exceeds maximum length") - case errors.Is(err, communitysuggestions.ErrDescriptionRequired): - writeError(w, http.StatusBadRequest, "InvalidRequest", "Suggestion description is required") - case errors.Is(err, communitysuggestions.ErrDescriptionTooLong): - writeError(w, http.StatusBadRequest, "InvalidRequest", "Suggestion description exceeds maximum length") - case errors.Is(err, communitysuggestions.ErrInvalidStatus): - writeError(w, http.StatusBadRequest, "InvalidStatus", "Invalid status value. Must be one of: open, under_review, approved, declined") - case errors.Is(err, communitysuggestions.ErrInvalidVoteValue): - writeError(w, http.StatusBadRequest, "InvalidVoteValue", "Invalid vote value. Must be 1 or -1") - case errors.Is(err, communitysuggestions.ErrInvalidSuggestionID): - writeError(w, http.StatusBadRequest, "InvalidRequest", "Invalid suggestion ID") - case errors.Is(err, communitysuggestions.ErrVoterRequired): - writeError(w, http.StatusBadRequest, "InvalidRequest", "Voter identification is required") - case errors.Is(err, communitysuggestions.ErrSubmitterRequired): - writeError(w, http.StatusBadRequest, "InvalidRequest", "Submitter identification is required") - default: - log.Printf("Unhandled validation error in community suggestion handler: %v", err) - writeError(w, http.StatusBadRequest, "InvalidRequest", "The request contains invalid data") - } - case communitysuggestions.IsNotFound(err): - writeError(w, http.StatusNotFound, "NotFound", "The requested resource was not found") - case communitysuggestions.IsRateLimitError(err): - writeError(w, http.StatusTooManyRequests, "RateLimitExceeded", "Too many suggestions. Please try again later") - case communitysuggestions.IsAuthorizationError(err): - writeError(w, http.StatusForbidden, "Forbidden", "You are not authorized to perform this action") - default: - log.Printf("XRPC community suggestion handler error: %v", err) - writeError(w, http.StatusInternalServerError, "InternalServerError", "An internal error occurred") - } + errorMapper.Write(w, err) } diff --git a/internal/api/handlers/discover/errors.go b/internal/api/handlers/discover/errors.go index 01c26f3..dd3d45d 100644 --- a/internal/api/handlers/discover/errors.go +++ b/internal/api/handlers/discover/errors.go @@ -1,43 +1,20 @@ package discover import ( - "Coves/internal/core/discover" - "encoding/json" - "errors" - "log" "net/http" -) - -// XRPCError represents an XRPC error response -type XRPCError struct { - Error string `json:"error"` - Message string `json:"message"` -} - -// writeError writes a JSON error response -func writeError(w http.ResponseWriter, status int, errorType, message string) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(status) - resp := XRPCError{ - Error: errorType, - Message: message, - } + "Coves/internal/api/xrpc" + "Coves/internal/core/discover" +) - if err := json.NewEncoder(w).Encode(resp); err != nil { - log.Printf("ERROR: Failed to encode error response: %v", err) - } -} +// errorMapper maps discover service errors to XRPC responses. +var errorMapper = xrpc.NewMapper("discover", + xrpc.MatchDetail(discover.IsValidationError, http.StatusBadRequest, "InvalidRequest"), + xrpc.Sentinel(discover.ErrInvalidCursor, http.StatusBadRequest, + "InvalidCursor", "The provided cursor is invalid"), +) -// handleServiceError maps service errors to HTTP responses +// handleServiceError maps service errors to HTTP responses. func handleServiceError(w http.ResponseWriter, err error) { - switch { - case discover.IsValidationError(err): - writeError(w, http.StatusBadRequest, "InvalidRequest", err.Error()) - case errors.Is(err, discover.ErrInvalidCursor): - writeError(w, http.StatusBadRequest, "InvalidCursor", "The provided cursor is invalid") - default: - log.Printf("ERROR: Discover service error: %v", err) - writeError(w, http.StatusInternalServerError, "InternalServerError", "An error occurred while fetching discover feed") - } + errorMapper.Write(w, err) } diff --git a/internal/api/handlers/errors.go b/internal/api/handlers/errors.go deleted file mode 100644 index e486817..0000000 --- a/internal/api/handlers/errors.go +++ /dev/null @@ -1,19 +0,0 @@ -package handlers - -import ( - "encoding/json" - "log" - "net/http" -) - -// WriteError writes a standardized JSON error response -func WriteError(w http.ResponseWriter, statusCode int, errorType, message string) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(statusCode) - if err := json.NewEncoder(w).Encode(map[string]interface{}{ - "error": errorType, - "message": message, - }); err != nil { - log.Printf("Failed to encode error response: %v", err) - } -} diff --git a/internal/api/handlers/post/delete.go b/internal/api/handlers/post/delete.go index a287928..46c7d2c 100644 --- a/internal/api/handlers/post/delete.go +++ b/internal/api/handlers/post/delete.go @@ -1,12 +1,13 @@ package post import ( - "Coves/internal/api/middleware" - "Coves/internal/core/posts" "encoding/json" - "errors" "log" "net/http" + + "Coves/internal/api/middleware" + "Coves/internal/api/xrpc" + "Coves/internal/core/posts" ) // DeleteHandler handles post deletion requests @@ -80,25 +81,23 @@ func (h *DeleteHandler) HandleDelete(w http.ResponseWriter, r *http.Request) { } } +// deleteErrorMapper narrows the package mapper for the delete path, which names +// the missing post and the refused action more precisely than the generic rules +// do; everything else it inherits. +// +// Note that deleting does NOT write to the caller's repo — posts live in the +// community's, and the delete authenticates with the community's service token. +// The service therefore strips the pds sentinels off a rejected community +// credential before it reaches here (see posts.communityCredentialFailure), so +// the inherited re-auth rule only ever fires on the caller's own session. +var deleteErrorMapper = errorMapper.With( + xrpc.Sentinel(posts.ErrNotFound, http.StatusNotFound, + "PostNotFound", "Post not found"), + xrpc.Sentinel(posts.ErrNotAuthorized, http.StatusForbidden, + "NotAuthorized", "You are not authorized to delete this post"), +) + // handleDeleteError maps delete-specific service errors to HTTP responses func handleDeleteError(w http.ResponseWriter, err error) { - switch { - case errors.Is(err, posts.ErrNotFound): - writeError(w, http.StatusNotFound, "PostNotFound", "Post not found") - - case errors.Is(err, posts.ErrNotAuthorized): - writeError(w, http.StatusForbidden, "NotAuthorized", "You are not authorized to delete this post") - - case errors.Is(err, posts.ErrCommunityNotFound): - writeError(w, http.StatusNotFound, "CommunityNotFound", "Community not found") - - case posts.IsValidationError(err): - writeError(w, http.StatusBadRequest, "InvalidRequest", err.Error()) - - default: - // Don't leak internal error details to clients - log.Printf("Unexpected error in post delete handler: %v", err) - writeError(w, http.StatusInternalServerError, "InternalServerError", - "An internal error occurred") - } + deleteErrorMapper.Write(w, err) } diff --git a/internal/api/handlers/post/errors.go b/internal/api/handlers/post/errors.go index c356aed..57bb45c 100644 --- a/internal/api/handlers/post/errors.go +++ b/internal/api/handlers/post/errors.go @@ -1,68 +1,51 @@ package post import ( + "net/http" + + "Coves/internal/api/xrpc" "Coves/internal/core/aggregators" "Coves/internal/core/posts" - "encoding/json" - "log" - "net/http" ) -type errorResponse struct { - Error string `json:"error"` - Message string `json:"message"` -} +// errorMapper maps post service errors to XRPC responses. +// +// Only the post-specific rules live here; dead sessions, other PDS failures, +// shared typed domain errors, and request-lifecycle errors come from xrpc's +// shared rules. Rules are tried in order, so the specific sentinels come before +// the broad predicates that also match them. +var errorMapper = xrpc.NewMapper("post", + // Ahead of posts.IsNotFound, which matches this sentinel too but answers + // with the less useful generic code. + xrpc.Sentinel(posts.ErrCommunityNotFound, http.StatusNotFound, + "CommunityNotFound", "Community not found"), + xrpc.Sentinel(posts.ErrNotAuthorized, http.StatusForbidden, + "NotAuthorized", "You are not authorized to post in this community"), + xrpc.Sentinel(posts.ErrBanned, http.StatusForbidden, + "Banned", "You are banned from this community"), + + // A content rule violation names the rule the post broke, which is the + // whole point of returning it — the client shows it to the author. + xrpc.As[*posts.ContentRuleViolation](http.StatusBadRequest, "ContentRuleViolation", + func(e *posts.ContentRuleViolation) string { return e.Error() }), + + xrpc.Sentinel(posts.ErrNotFound, http.StatusNotFound, + "NotFound", "Post not found"), + + xrpc.Match(aggregators.IsUnauthorized, http.StatusForbidden, + "NotAuthorized", "Aggregator not authorized to post in this community"), + xrpc.Match(aggregators.IsRateLimited, http.StatusTooManyRequests, + "RateLimitExceeded", "Rate limit exceeded. Please try again later."), + xrpc.Sentinel(posts.ErrRateLimitExceeded, http.StatusTooManyRequests, + "RateLimitExceeded", "Rate limit exceeded. Please try again later."), +) -// writeError writes a JSON error response +// writeError writes a JSON error response. func writeError(w http.ResponseWriter, statusCode int, errorType, message string) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(statusCode) - if err := json.NewEncoder(w).Encode(errorResponse{ - Error: errorType, - Message: message, - }); err != nil { - log.Printf("Failed to encode error response: %v", err) - } + xrpc.WriteError(w, statusCode, errorType, message) } -// handleServiceError maps service errors to HTTP responses +// handleServiceError maps service errors to HTTP responses. func handleServiceError(w http.ResponseWriter, err error) { - switch { - case err == posts.ErrCommunityNotFound: - writeError(w, http.StatusNotFound, "CommunityNotFound", - "Community not found") - - case err == posts.ErrNotAuthorized: - writeError(w, http.StatusForbidden, "NotAuthorized", - "You are not authorized to post in this community") - - case err == posts.ErrBanned: - writeError(w, http.StatusForbidden, "Banned", - "You are banned from this community") - - case posts.IsContentRuleViolation(err): - writeError(w, http.StatusBadRequest, "ContentRuleViolation", err.Error()) - - case posts.IsValidationError(err): - writeError(w, http.StatusBadRequest, "InvalidRequest", err.Error()) - - case posts.IsNotFound(err): - writeError(w, http.StatusNotFound, "NotFound", err.Error()) - - // Check aggregator authorization errors - case aggregators.IsUnauthorized(err): - writeError(w, http.StatusForbidden, "NotAuthorized", - "Aggregator not authorized to post in this community") - - // Check both aggregator and post rate limit errors - case aggregators.IsRateLimited(err) || err == posts.ErrRateLimitExceeded: - writeError(w, http.StatusTooManyRequests, "RateLimitExceeded", - "Rate limit exceeded. Please try again later.") - - default: - // Don't leak internal error details to clients - log.Printf("Unexpected error in post handler: %v", err) - writeError(w, http.StatusInternalServerError, "InternalServerError", - "An internal error occurred") - } + errorMapper.Write(w, err) } diff --git a/internal/api/handlers/post/errors_test.go b/internal/api/handlers/post/errors_test.go new file mode 100644 index 0000000..145972c --- /dev/null +++ b/internal/api/handlers/post/errors_test.go @@ -0,0 +1,187 @@ +package post + +import ( + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "Coves/internal/api/xrpc" + "Coves/internal/atproto/pds" + "Coves/internal/core/posts" +) + +// Deleting a post can fail with any PDS error. This package mapped none of +// them, so an expired session answered 500 and clients never learned to +// re-authenticate. +// +// The errors here are the ones raised against the CALLER's session. Posts live +// in the community's repo, so a failure of the community's own service +// credentials must not land in this set — see +// TestCommunityCredentialFailureIsNot401. +func TestDeleteMapsPDSErrors(t *testing.T) { + tests := []struct { + name string + err error + wantStatus int + wantCode string + }{ + { + name: "expired token", + err: fmt.Errorf("DeleteRecord: %w: expired", pds.ErrUnauthorized), + wantStatus: http.StatusUnauthorized, + wantCode: "AuthRequired", + }, + { + name: "session could not be resumed", + err: fmt.Errorf("failed to create PDS client: %w", + fmt.Errorf("resume: %w: revoked", pds.ErrSessionExpired)), + wantStatus: http.StatusUnauthorized, + wantCode: "AuthRequired", + }, + { + name: "missing scope", + err: fmt.Errorf("DeleteRecord: %w", pds.ErrForbidden), + wantStatus: http.StatusForbidden, + wantCode: "PermissionDenied", + }, + { + name: "pds rate limit", + err: fmt.Errorf("DeleteRecord: %w", pds.ErrRateLimited), + wantStatus: http.StatusTooManyRequests, + wantCode: "RateLimitExceeded", + }, + { + // Delete keeps its own more specific codes. + name: "missing post", + err: fmt.Errorf("service: %w", posts.ErrNotFound), + wantStatus: http.StatusNotFound, + wantCode: "PostNotFound", + }, + { + name: "not the author", + err: fmt.Errorf("service: %w", posts.ErrNotAuthorized), + wantStatus: http.StatusForbidden, + wantCode: "NotAuthorized", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + rec := httptest.NewRecorder() + handleDeleteError(rec, tt.err) + assertXRPCError(t, rec, tt.wantStatus, tt.wantCode) + }) + } +} + +// TestCommunityCredentialCannotTriggerReauth guards the one place reauth-first +// could misfire. Posts live in the community's repo, so a delete authenticates +// with the community's service token; if a rejection of THAT credential reached +// the mapper carrying pds.ErrUnauthorized, the boundary would read it as the +// caller's session being dead and answer 401 — telling a user with a healthy +// session to sign in over a server-side problem, and hiding the outage from 5xx. +// +// posts.DeletePost strips the sentinel before returning, so the error arrives +// unclassified. This asserts the boundary behaviour that depends on it. +func TestCommunityCredentialCannotTriggerReauth(t *testing.T) { + // Shaped as posts.communityCredentialFailure builds it: %v, not %w. + err := fmt.Errorf("community PDS credentials rejected during delete post for %s: %v", + "did:plc:community", fmt.Errorf("DeleteRecord: %w: bad token", pds.ErrUnauthorized)) + + rec := httptest.NewRecorder() + handleDeleteError(rec, err) + + body := assertXRPCError(t, rec, http.StatusInternalServerError, "InternalServerError") + if rec.Code == http.StatusUnauthorized { + t.Fatal("a community credential failure must never tell the caller to re-authenticate") + } + if strings.Contains(body.Message, "did:plc:community") { + t.Errorf("internal detail leaked: %q", body.Message) + } +} + +// The post error codes are a client contract; this pins them. Note that the +// wrapped forms below newly reach these codes at all: the switch this replaced +// compared with ==, so a wrapped sentinel fell through to 500. +func TestPostErrorCodes(t *testing.T) { + tests := []struct { + err error + wantStatus int + wantCode string + }{ + {posts.ErrCommunityNotFound, http.StatusNotFound, "CommunityNotFound"}, + {posts.ErrNotAuthorized, http.StatusForbidden, "NotAuthorized"}, + {posts.ErrBanned, http.StatusForbidden, "Banned"}, + {posts.ErrNotFound, http.StatusNotFound, "NotFound"}, + {posts.ErrRateLimitExceeded, http.StatusTooManyRequests, "RateLimitExceeded"}, + } + + for _, tt := range tests { + t.Run(tt.wantCode, func(t *testing.T) { + rec := httptest.NewRecorder() + handleServiceError(rec, fmt.Errorf("service: %w", tt.err)) + assertXRPCError(t, rec, tt.wantStatus, tt.wantCode) + }) + } +} + +// posts.ErrCommunityNotFound must keep beating the generic not-found rule that +// also matches it. +func TestCommunityNotFoundBeatsGenericNotFound(t *testing.T) { + rec := httptest.NewRecorder() + handleServiceError(rec, fmt.Errorf("service: %w", posts.ErrCommunityNotFound)) + assertXRPCError(t, rec, http.StatusNotFound, "CommunityNotFound") +} + +// Typed errors carry their own client-facing text; wrapper context added by the +// service on the way up must not ride along into the response. +func TestTypedPostErrorsUseTheirOwnMessage(t *testing.T) { + rec := httptest.NewRecorder() + handleServiceError(rec, fmt.Errorf("createPost: %w", + posts.NewContentRuleViolation("requireText", "posts must have text"))) + body := assertXRPCError(t, rec, http.StatusBadRequest, "ContentRuleViolation") + if strings.Contains(body.Message, "createPost") { + t.Errorf("wrapper context leaked into the client message: %q", body.Message) + } + + rec = httptest.NewRecorder() + handleServiceError(rec, fmt.Errorf("createPost: %w", posts.NewValidationError("uri", "is malformed"))) + body = assertXRPCError(t, rec, http.StatusBadRequest, "InvalidRequest") + if body.Message != "uri: is malformed" { + t.Errorf("message = %q, want just the typed error's text", body.Message) + } +} + +// The sentinel rules replaced err.Error() echoes with fixed strings, so an +// internal wrapper can no longer reach the client. +func TestSentinelMessagesAreFixed(t *testing.T) { + rec := httptest.NewRecorder() + handleServiceError(rec, fmt.Errorf("repo query against 10.0.0.4 failed: %w", posts.ErrNotFound)) + + body := assertXRPCError(t, rec, http.StatusNotFound, "NotFound") + if body.Message != "Post not found" { + t.Errorf("message = %q, want the fixed string", body.Message) + } + if strings.Contains(body.Message, "10.0.0.4") { + t.Errorf("internal detail leaked: %q", body.Message) + } +} + +func assertXRPCError(t *testing.T, rec *httptest.ResponseRecorder, wantStatus int, wantCode string) xrpc.Error { + t.Helper() + + if rec.Code != wantStatus { + t.Errorf("status = %d, want %d", rec.Code, wantStatus) + } + var body xrpc.Error + if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { + t.Fatalf("body is not valid JSON: %v", err) + } + if body.Error != wantCode { + t.Errorf("code = %q, want %q", body.Error, wantCode) + } + return body +} diff --git a/internal/api/handlers/timeline/errors.go b/internal/api/handlers/timeline/errors.go index c91f039..a3c7bb6 100644 --- a/internal/api/handlers/timeline/errors.go +++ b/internal/api/handlers/timeline/errors.go @@ -1,45 +1,31 @@ package timeline import ( - "Coves/internal/core/timeline" - "encoding/json" - "errors" - "log" "net/http" + + "Coves/internal/api/xrpc" + "Coves/internal/core/timeline" ) -// XRPCError represents an XRPC error response -type XRPCError struct { - Error string `json:"error"` - Message string `json:"message"` -} +// errorMapper maps timeline service errors to XRPC responses. +// +// Rule order matches the switch this replaced: a cursor problem reported as a +// typed validation error answers InvalidRequest, and only a bare +// ErrInvalidCursor answers InvalidCursor. +var errorMapper = xrpc.NewMapper("timeline", + xrpc.MatchDetail(timeline.IsValidationError, http.StatusBadRequest, "InvalidRequest"), + xrpc.Sentinel(timeline.ErrInvalidCursor, http.StatusBadRequest, + "InvalidCursor", "The provided cursor is invalid"), + xrpc.Sentinel(timeline.ErrUnauthorized, http.StatusUnauthorized, + "AuthenticationRequired", "User must be authenticated"), +) -// writeError writes a JSON error response +// writeError writes a JSON error response. func writeError(w http.ResponseWriter, status int, errorType, message string) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(status) - - resp := XRPCError{ - Error: errorType, - Message: message, - } - - if err := json.NewEncoder(w).Encode(resp); err != nil { - log.Printf("ERROR: Failed to encode error response: %v", err) - } + xrpc.WriteError(w, status, errorType, message) } -// handleServiceError maps service errors to HTTP responses +// handleServiceError maps service errors to HTTP responses. func handleServiceError(w http.ResponseWriter, err error) { - switch { - case timeline.IsValidationError(err): - writeError(w, http.StatusBadRequest, "InvalidRequest", err.Error()) - case errors.Is(err, timeline.ErrInvalidCursor): - writeError(w, http.StatusBadRequest, "InvalidCursor", "The provided cursor is invalid") - case errors.Is(err, timeline.ErrUnauthorized): - writeError(w, http.StatusUnauthorized, "AuthenticationRequired", "User must be authenticated") - default: - log.Printf("ERROR: Timeline service error: %v", err) - writeError(w, http.StatusInternalServerError, "InternalServerError", "An error occurred while fetching timeline") - } + errorMapper.Write(w, err) } diff --git a/internal/api/handlers/user/delete.go b/internal/api/handlers/user/delete.go index d054312..53963c5 100644 --- a/internal/api/handlers/user/delete.go +++ b/internal/api/handlers/user/delete.go @@ -88,64 +88,41 @@ func (h *DeleteHandler) HandleDeleteAccount(w http.ResponseWriter, r *http.Reque } } -// writeJSONError writes a JSON error response -// Marshals JSON before writing headers to catch encoding errors -func writeJSONError(w http.ResponseWriter, statusCode int, errorType, message string) { - responseBytes, err := json.Marshal(map[string]interface{}{ - "error": errorType, - "message": message, - }) - if err != nil { - // Fallback to plain text if JSON encoding fails (should never happen with simple strings) - slog.Error("failed to marshal error response", slog.String("error", err.Error())) - w.Header().Set("Content-Type", "text/plain") - w.WriteHeader(statusCode) - _, _ = w.Write([]byte(message)) - return - } - - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(statusCode) - if _, writeErr := w.Write(responseBytes); writeErr != nil { - slog.Warn("failed to write error response", slog.String("error", writeErr.Error())) - } -} - // handleServiceError maps service errors to HTTP responses. // operation is a human-readable label for log messages (e.g. "account deletion", "get profile"). func handleServiceError(w http.ResponseWriter, err error, userDID, operation string) { - // Check for specific error types - switch { - case errors.Is(err, users.ErrUserNotFound): - writeJSONError(w, http.StatusNotFound, "AccountNotFound", "Account not found") + if err == nil { + slog.Error(operation+" reached its error path without an error", + slog.String("did", userDID)) + accountErrorMapper.Write(w, err) + return + } - case errors.Is(err, context.DeadlineExceeded): - slog.Error(operation+" timed out", + // The DID and operation are worth recording on every failure here, not just + // the unmapped ones the mapper logs, so account problems can be traced to a + // user. Severity follows the answer, though: a mistyped DID or a missing + // account is the caller's ordinary mistake, and logging those at ERROR — as + // this did before the outcome was available to branch on — buries real + // faults in alert noise. + mapping, matched := accountErrorMapper.Resolve(err) + switch { + case errors.Is(err, context.Canceled): + slog.Info(operation+" canceled", slog.String("did", userDID), slog.String("error", err.Error()), ) - writeJSONError(w, http.StatusGatewayTimeout, "Timeout", "Request timed out") - - case errors.Is(err, context.Canceled): - slog.Info(operation+" canceled", + case matched && mapping.Status < http.StatusInternalServerError: + slog.Info(operation+" rejected", slog.String("did", userDID), + slog.String("code", mapping.Code), slog.String("error", err.Error()), ) - writeJSONError(w, http.StatusBadRequest, "RequestCanceled", "Request was canceled") - default: - // Check for InvalidDIDError - var invalidDIDErr *users.InvalidDIDError - if errors.As(err, &invalidDIDErr) { - writeJSONError(w, http.StatusBadRequest, "InvalidDID", invalidDIDErr.Error()) - return - } - - // Internal server error - don't leak details slog.Error(operation+" failed", slog.String("did", userDID), slog.String("error", err.Error()), ) - writeJSONError(w, http.StatusInternalServerError, "InternalServerError", "An internal error occurred") } + + accountErrorMapper.Write(w, err) } diff --git a/internal/api/handlers/user/errors.go b/internal/api/handlers/user/errors.go new file mode 100644 index 0000000..d14ff2d --- /dev/null +++ b/internal/api/handlers/user/errors.go @@ -0,0 +1,84 @@ +package user + +import ( + "net/http" + + "Coves/internal/api/xrpc" + "Coves/internal/atproto/pds" + "Coves/internal/core/users" +) + +// writeJSONError writes a JSON error response. +func writeJSONError(w http.ResponseWriter, statusCode int, errorType, message string) { + xrpc.WriteError(w, statusCode, errorType, message) +} + +// writeUpdateProfileError writes a JSON error response for update profile failures. +func writeUpdateProfileError(w http.ResponseWriter, statusCode int, errorType, message string) { + xrpc.WriteError(w, statusCode, errorType, message) +} + +// accountErrorMapper maps users service errors for the account endpoints. +// +// The users service talks to the PDS admin API over plain HTTP rather than +// through internal/atproto/pds, so it raises no PDS sentinels; the inherited +// rules cost nothing and cover it if that changes. +var accountErrorMapper = xrpc.NewMapper("user", + xrpc.Sentinel(users.ErrUserNotFound, http.StatusNotFound, + "AccountNotFound", "Account not found"), + xrpc.As[*users.InvalidDIDError](http.StatusBadRequest, "InvalidDID", + func(e *users.InvalidDIDError) string { return e.Error() }), +) + +// updateProfileMapper is the base for the updateProfile endpoint, the only +// handler that drives a pds.Client directly. +// +// It answers 401 with AuthExpired rather than the default AuthRequired because +// clients already key on that code here. 403 is left to the inherited rule: it +// is a permissions problem — an OAuth grant predating the blob:*/* scope, say — +// not an expired session, so it must not trigger a sign-out, even though +// signing in again is what re-grants the scope. +var updateProfileMapper = xrpc.NewMapper("user.updateProfile"). + WithReauth("AuthExpired", "Your session may have expired. Please re-authenticate.") + +// sessionRestoreMapper covers building the PDS client. If that fails at all +// there is no usable session for this request, whatever the cause, so every +// outcome is 401 — the fallback included. The cause is still logged. +var sessionRestoreMapper = xrpc.NewMapper("user.updateProfile.session"). + WithReauth("SessionError", "Failed to restore session. Please sign in again."). + WithFallback(http.StatusUnauthorized, "SessionError", "Failed to restore session. Please sign in again.") + +// newBlobUploadMapper is shared by the avatar and banner uploads, which differ +// only in the code they use for an oversized image. +func newBlobUploadMapper(tooLargeCode, tooLargeMessage, failureMessage string) *xrpc.Mapper { + return updateProfileMapper. + WithFallback(http.StatusInternalServerError, "BlobUploadFailed", failureMessage). + With( + xrpc.Sentinel(pds.ErrForbidden, http.StatusForbidden, "PermissionDenied", + "Your session does not have permission to upload images. Sign out and back in to grant it."), + xrpc.Sentinel(pds.ErrRateLimited, http.StatusTooManyRequests, "RateLimited", + "Too many requests. Please try again later."), + xrpc.Sentinel(pds.ErrPayloadTooLarge, http.StatusRequestEntityTooLarge, + tooLargeCode, tooLargeMessage), + ) +} + +var ( + avatarUploadMapper = newBlobUploadMapper( + "AvatarTooLarge", "Avatar exceeds PDS size limit.", "Failed to upload avatar") + bannerUploadMapper = newBlobUploadMapper( + "BannerTooLarge", "Banner exceeds PDS size limit.", "Failed to upload banner") + + // putProfileMapper covers writing the profile record itself. + putProfileMapper = updateProfileMapper. + WithFallback(http.StatusInternalServerError, "PDSError", "Failed to update profile"). + With( + xrpc.Sentinel(pds.ErrForbidden, http.StatusForbidden, "PermissionDenied", + "Your session does not have permission to update your profile. Sign out and back in to grant it."), + xrpc.Sentinel(pds.ErrRateLimited, http.StatusTooManyRequests, "RateLimited", + "Too many requests. Please try again later."), + // Same code as the inherited rule, kept for its profile-specific wording. + xrpc.Sentinel(pds.ErrPayloadTooLarge, http.StatusRequestEntityTooLarge, + "PayloadTooLarge", "Profile data exceeds PDS size limit."), + ) +) diff --git a/internal/api/handlers/user/update_profile.go b/internal/api/handlers/user/update_profile.go index 82995e7..4bd60ab 100644 --- a/internal/api/handlers/user/update_profile.go +++ b/internal/api/handlers/user/update_profile.go @@ -3,7 +3,6 @@ package user import ( "context" "encoding/json" - "errors" "fmt" "log/slog" "net/http" @@ -201,8 +200,7 @@ func (h *UpdateProfileHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) slog.String("did", userDID), slog.String("error", err.Error()), ) - writeUpdateProfileError(w, http.StatusUnauthorized, "SessionError", - "Failed to restore session. Please sign in again.") + sessionRestoreMapper.Write(w, err) return } @@ -229,22 +227,7 @@ func (h *UpdateProfileHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) slog.String("did", userDID), slog.String("error", err.Error()), ) - // Map specific PDS errors to user-friendly messages. - // 403 is a permissions problem (e.g. the OAuth grant predates the - // blob:*/* scope), NOT an expired session — it must not trigger a - // client sign-out, and signing in again is what re-grants the scope. - switch { - case errors.Is(err, pds.ErrForbidden): - writeUpdateProfileError(w, http.StatusForbidden, "PermissionDenied", "Your session does not have permission to upload images. Sign out and back in to grant it.") - case errors.Is(err, pds.ErrUnauthorized): - writeUpdateProfileError(w, http.StatusUnauthorized, "AuthExpired", "Your session may have expired. Please re-authenticate.") - case errors.Is(err, pds.ErrRateLimited): - writeUpdateProfileError(w, http.StatusTooManyRequests, "RateLimited", "Too many requests. Please try again later.") - case errors.Is(err, pds.ErrPayloadTooLarge): - writeUpdateProfileError(w, http.StatusRequestEntityTooLarge, "AvatarTooLarge", "Avatar exceeds PDS size limit.") - default: - writeUpdateProfileError(w, http.StatusInternalServerError, "BlobUploadFailed", "Failed to upload avatar") - } + avatarUploadMapper.Write(w, err) return } if avatarRef == nil || avatarRef.Ref == nil || avatarRef.Type == "" { @@ -268,20 +251,7 @@ func (h *UpdateProfileHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) slog.String("did", userDID), slog.String("error", err.Error()), ) - // Map specific PDS errors to user-friendly messages. - // 403 is a permissions problem, not an expired session (see avatar case). - switch { - case errors.Is(err, pds.ErrForbidden): - writeUpdateProfileError(w, http.StatusForbidden, "PermissionDenied", "Your session does not have permission to upload images. Sign out and back in to grant it.") - case errors.Is(err, pds.ErrUnauthorized): - writeUpdateProfileError(w, http.StatusUnauthorized, "AuthExpired", "Your session may have expired. Please re-authenticate.") - case errors.Is(err, pds.ErrRateLimited): - writeUpdateProfileError(w, http.StatusTooManyRequests, "RateLimited", "Too many requests. Please try again later.") - case errors.Is(err, pds.ErrPayloadTooLarge): - writeUpdateProfileError(w, http.StatusRequestEntityTooLarge, "BannerTooLarge", "Banner exceeds PDS size limit.") - default: - writeUpdateProfileError(w, http.StatusInternalServerError, "BlobUploadFailed", "Failed to upload banner") - } + bannerUploadMapper.Write(w, err) return } if bannerRef == nil || bannerRef.Ref == nil || bannerRef.Type == "" { @@ -305,20 +275,7 @@ func (h *UpdateProfileHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) slog.String("pds_url", session.HostURL), slog.String("error", err.Error()), ) - // Map PDS errors to user-friendly messages. - // 403 is a permissions problem, not an expired session (see avatar case). - switch { - case errors.Is(err, pds.ErrForbidden): - writeUpdateProfileError(w, http.StatusForbidden, "PermissionDenied", "Your session does not have permission to update your profile. Sign out and back in to grant it.") - case errors.Is(err, pds.ErrUnauthorized): - writeUpdateProfileError(w, http.StatusUnauthorized, "AuthExpired", "Your session may have expired. Please re-authenticate.") - case errors.Is(err, pds.ErrRateLimited): - writeUpdateProfileError(w, http.StatusTooManyRequests, "RateLimited", "Too many requests. Please try again later.") - case errors.Is(err, pds.ErrPayloadTooLarge): - writeUpdateProfileError(w, http.StatusRequestEntityTooLarge, "PayloadTooLarge", "Profile data exceeds PDS size limit.") - default: - writeUpdateProfileError(w, http.StatusInternalServerError, "PDSError", "Failed to update profile") - } + putProfileMapper.Write(w, err) return } @@ -355,25 +312,3 @@ func isValidImageMimeType(mimeType string) bool { return false } } - -// writeUpdateProfileError writes a JSON error response for update profile failures -func writeUpdateProfileError(w http.ResponseWriter, statusCode int, errorType, message string) { - responseBytes, err := json.Marshal(map[string]interface{}{ - "error": errorType, - "message": message, - }) - if err != nil { - // Fallback to plain text if JSON encoding fails - slog.Error("failed to marshal error response", slog.String("error", err.Error())) - w.Header().Set("Content-Type", "text/plain") - w.WriteHeader(statusCode) - _, _ = w.Write([]byte(message)) - return - } - - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(statusCode) - if _, writeErr := w.Write(responseBytes); writeErr != nil { - slog.Warn("failed to write error response", slog.String("error", writeErr.Error())) - } -} diff --git a/internal/api/handlers/userblock/errors.go b/internal/api/handlers/userblock/errors.go index 7fcd658..c832d13 100644 --- a/internal/api/handlers/userblock/errors.go +++ b/internal/api/handlers/userblock/errors.go @@ -1,59 +1,29 @@ package userblock import ( - "Coves/internal/atproto/pds" - "Coves/internal/core/userblocks" - "encoding/json" - "errors" - "log/slog" "net/http" + + "Coves/internal/api/xrpc" + "Coves/internal/core/userblocks" ) -// XRPCError represents an XRPC error response -type XRPCError struct { - Error string `json:"error"` - Message string `json:"message"` -} +// errorMapper maps user block service errors to XRPC responses. +// +// This package was the only one that mapped the full set of PDS errors by +// hand; those rules now live in xrpc and every handler package gets them. +var errorMapper = xrpc.NewMapper("userblock", + xrpc.SentinelDetail(userblocks.ErrBlockNotFound, http.StatusNotFound, "NotFound"), + xrpc.SentinelDetail(userblocks.ErrBlockAlreadyExists, http.StatusConflict, "AlreadyExists"), + xrpc.Sentinel(userblocks.ErrCannotBlockSelf, http.StatusBadRequest, + "InvalidRequest", "cannot block yourself"), +) -// writeError writes an XRPC error response +// writeError writes an XRPC error response. func writeError(w http.ResponseWriter, status int, errCode, message string) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(status) - if err := json.NewEncoder(w).Encode(XRPCError{ - Error: errCode, - Message: message, - }); err != nil { - slog.Error("Failed to encode error response", "error", err) - } + xrpc.WriteError(w, status, errCode, message) } -// handleServiceError converts user block service errors to appropriate HTTP responses +// handleServiceError converts user block service errors to appropriate HTTP responses. func handleServiceError(w http.ResponseWriter, err error) { - switch { - case errors.Is(err, userblocks.ErrBlockNotFound): - writeError(w, http.StatusNotFound, "NotFound", err.Error()) - case errors.Is(err, userblocks.ErrBlockAlreadyExists): - writeError(w, http.StatusConflict, "AlreadyExists", err.Error()) - case errors.Is(err, userblocks.ErrCannotBlockSelf): - writeError(w, http.StatusBadRequest, "InvalidRequest", "cannot block yourself") - // PDS-specific errors. 403 is a permissions problem (e.g. missing OAuth - // scope), not an expired session — it must not trigger a client sign-out. - case errors.Is(err, pds.ErrForbidden): - writeError(w, http.StatusForbidden, "PermissionDenied", "Your session does not have permission for this action. Sign out and back in to grant it.") - case errors.Is(err, pds.ErrUnauthorized): - writeError(w, http.StatusUnauthorized, "AuthRequired", "Authentication required or session expired") - case errors.Is(err, pds.ErrBadRequest): - writeError(w, http.StatusBadRequest, "InvalidRequest", "Invalid request to PDS") - case errors.Is(err, pds.ErrNotFound): - writeError(w, http.StatusNotFound, "NotFound", "Record not found on PDS") - case errors.Is(err, pds.ErrConflict): - writeError(w, http.StatusConflict, "Conflict", "Record was modified by another operation") - case errors.Is(err, pds.ErrRateLimited): - writeError(w, http.StatusTooManyRequests, "RateLimitExceeded", "Too many requests, please try again later") - case errors.Is(err, pds.ErrPayloadTooLarge): - writeError(w, http.StatusRequestEntityTooLarge, "PayloadTooLarge", "Request payload exceeds size limit") - default: - slog.Error("XRPC user block handler error", "error", err) - writeError(w, http.StatusInternalServerError, "InternalServerError", "An internal error occurred") - } + errorMapper.Write(w, err) } diff --git a/internal/api/handlers/vote/errors.go b/internal/api/handlers/vote/errors.go index f34ae6c..b8d0807 100644 --- a/internal/api/handlers/vote/errors.go +++ b/internal/api/handlers/vote/errors.go @@ -1,54 +1,49 @@ package vote import ( - "Coves/internal/core/votes" - "encoding/json" - "errors" - "log" "net/http" + + "Coves/internal/api/xrpc" + "Coves/internal/core/votes" ) -// XRPCError represents an XRPC error response -type XRPCError struct { - Error string `json:"error"` - Message string `json:"message"` -} +// XRPCError represents an XRPC error response. +type XRPCError = xrpc.Error + +// errorMapper maps vote service errors to XRPC responses. +// +// Only the vote-specific rules live here. Dead sessions, other PDS failures, +// shared typed domain errors, and request-lifecycle errors come from xrpc's +// shared rules, so this package cannot fall behind on them the way a +// hand-written switch did. +// +// Error names are part of the client contract: keep them UpperCamelCase and +// stable. +var errorMapper = xrpc.NewMapper("vote", + // Matches: social.coves.feed.vote.delete#VoteNotFound + xrpc.Sentinel(votes.ErrVoteNotFound, http.StatusNotFound, + "VoteNotFound", "No vote found for this subject"), + xrpc.Sentinel(votes.ErrInvalidDirection, http.StatusBadRequest, + "InvalidRequest", "Vote direction must be 'up' or 'down'"), + // Matches: social.coves.feed.vote.create#InvalidSubject + xrpc.Sentinel(votes.ErrInvalidSubject, http.StatusBadRequest, + "InvalidSubject", "The subject reference is invalid or malformed"), + xrpc.Sentinel(votes.ErrVoteAlreadyExists, http.StatusConflict, + "AlreadyExists", "Vote already exists"), + // Matches: social.coves.feed.vote.create#NotAuthorized, + // social.coves.feed.vote.delete#NotAuthorized + xrpc.Sentinel(votes.ErrNotAuthorized, http.StatusForbidden, + "NotAuthorized", "User is not authorized to vote on this content"), + xrpc.Sentinel(votes.ErrBanned, http.StatusForbidden, + "NotAuthorized", "User is not authorized to vote on this content"), +) -// writeError writes an XRPC error response -func writeError(w http.ResponseWriter, status int, error, message string) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(status) - if err := json.NewEncoder(w).Encode(XRPCError{ - Error: error, - Message: message, - }); err != nil { - log.Printf("Failed to encode error response: %v", err) - } +// writeError writes an XRPC error response. +func writeError(w http.ResponseWriter, status int, code, message string) { + xrpc.WriteError(w, status, code, message) } -// handleServiceError converts service errors to appropriate HTTP responses -// Error names MUST match lexicon definitions exactly (UpperCamelCase) -// Uses errors.Is() to handle wrapped errors correctly +// handleServiceError converts service errors to appropriate HTTP responses. func handleServiceError(w http.ResponseWriter, err error) { - switch { - case errors.Is(err, votes.ErrVoteNotFound): - // Matches: social.coves.feed.vote.delete#VoteNotFound - writeError(w, http.StatusNotFound, "VoteNotFound", "No vote found for this subject") - case errors.Is(err, votes.ErrInvalidDirection): - writeError(w, http.StatusBadRequest, "InvalidRequest", "Vote direction must be 'up' or 'down'") - case errors.Is(err, votes.ErrInvalidSubject): - // Matches: social.coves.feed.vote.create#InvalidSubject - writeError(w, http.StatusBadRequest, "InvalidSubject", "The subject reference is invalid or malformed") - case errors.Is(err, votes.ErrVoteAlreadyExists): - writeError(w, http.StatusConflict, "AlreadyExists", "Vote already exists") - case errors.Is(err, votes.ErrNotAuthorized): - // Matches: social.coves.feed.vote.create#NotAuthorized, social.coves.feed.vote.delete#NotAuthorized - writeError(w, http.StatusForbidden, "NotAuthorized", "User is not authorized to vote on this content") - case errors.Is(err, votes.ErrBanned): - writeError(w, http.StatusForbidden, "NotAuthorized", "User is not authorized to vote on this content") - default: - // Internal server error - log the actual error for debugging - log.Printf("XRPC handler error: %v", err) - writeError(w, http.StatusInternalServerError, "InternalServerError", "An internal error occurred") - } + errorMapper.Write(w, err) } diff --git a/internal/api/handlers/vote/errors_test.go b/internal/api/handlers/vote/errors_test.go new file mode 100644 index 0000000..de67999 --- /dev/null +++ b/internal/api/handlers/vote/errors_test.go @@ -0,0 +1,119 @@ +package vote + +import ( + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "testing" + + "Coves/internal/atproto/pds" + "Coves/internal/core/votes" +) + +// TestExpiredSessionAnswers401 covers the reported bug: voting with an expired +// OAuth session answered 500 InternalServerError, so clients had no signal to +// re-authenticate and users saw "an internal error occurred" until they signed +// out by hand. +// +// The service reports its own ErrNotAuthorized with the PDS cause attached, and +// the mapper has to see past the former to the latter. +func TestExpiredSessionAnswers401(t *testing.T) { + tests := []struct { + name string + err error + wantStatus int + wantCode string + }{ + { + name: "pds rejected the token", + err: fmt.Errorf("%w: %w", votes.ErrNotAuthorized, + fmt.Errorf("CreateRecord: %w: expired token", pds.ErrUnauthorized)), + wantStatus: http.StatusUnauthorized, + wantCode: "AuthRequired", + }, + { + name: "session could not be resumed", + err: fmt.Errorf("failed to create PDS client: %w", + fmt.Errorf("failed to resume OAuth session: %w: revoked", pds.ErrSessionExpired)), + wantStatus: http.StatusUnauthorized, + wantCode: "AuthRequired", + }, + { + // A scope gap is not a dead session and must not sign the user out. + name: "missing scope stays 403", + err: fmt.Errorf("%w: %w", votes.ErrNotAuthorized, + fmt.Errorf("CreateRecord: %w", pds.ErrForbidden)), + wantStatus: http.StatusForbidden, + wantCode: "NotAuthorized", + }, + { + // An appview-level refusal has no PDS cause and keeps its own answer. + name: "domain refusal stays 403", + err: votes.ErrNotAuthorized, + wantStatus: http.StatusForbidden, + wantCode: "NotAuthorized", + }, + { + name: "pds rate limit is no longer a 500", + err: fmt.Errorf("CreateRecord: %w", pds.ErrRateLimited), + wantStatus: http.StatusTooManyRequests, + wantCode: "RateLimitExceeded", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + rec := httptest.NewRecorder() + handleServiceError(rec, tt.err) + + if rec.Code != tt.wantStatus { + t.Errorf("status = %d, want %d", rec.Code, tt.wantStatus) + } + var body XRPCError + if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { + t.Fatalf("body is not valid JSON: %v", err) + } + if body.Error != tt.wantCode { + t.Errorf("code = %q, want %q", body.Error, tt.wantCode) + } + }) + } +} + +// The vote error codes are a client contract; this pins them so a mapper +// refactor cannot quietly rename one. The wrapped forms are the ones that +// changed: errors.Is now matches where == did not. +func TestVoteErrorCodes(t *testing.T) { + tests := []struct { + err error + wantStatus int + wantCode string + }{ + {votes.ErrVoteNotFound, http.StatusNotFound, "VoteNotFound"}, + {votes.ErrInvalidDirection, http.StatusBadRequest, "InvalidRequest"}, + {votes.ErrInvalidSubject, http.StatusBadRequest, "InvalidSubject"}, + {votes.ErrVoteAlreadyExists, http.StatusConflict, "AlreadyExists"}, + {votes.ErrNotAuthorized, http.StatusForbidden, "NotAuthorized"}, + {votes.ErrBanned, http.StatusForbidden, "NotAuthorized"}, + } + + for _, tt := range tests { + t.Run(tt.wantCode+"/"+tt.err.Error(), func(t *testing.T) { + rec := httptest.NewRecorder() + // Wrapped, because service layers add context on the way up. + handleServiceError(rec, fmt.Errorf("service: %w", tt.err)) + + if rec.Code != tt.wantStatus { + t.Errorf("status = %d, want %d", rec.Code, tt.wantStatus) + } + var body XRPCError + if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { + t.Fatalf("body is not valid JSON: %v", err) + } + if body.Error != tt.wantCode { + t.Errorf("code = %q, want %q", body.Error, tt.wantCode) + } + }) + } +} diff --git a/internal/api/routes/user.go b/internal/api/routes/user.go index faddbaa..1f0ecbf 100644 --- a/internal/api/routes/user.go +++ b/internal/api/routes/user.go @@ -3,6 +3,7 @@ package routes import ( "Coves/internal/api/handlers/user" "Coves/internal/api/middleware" + "Coves/internal/api/xrpc" "Coves/internal/core/userblocks" "Coves/internal/core/users" "encoding/json" @@ -127,7 +128,19 @@ func (h *UserHandler) GetProfile(w http.ResponseWriter, r *http.Request) { // Resolve handle to DID resolvedDID, err := h.userService.ResolveHandleToDID(ctx, actor) if err != nil { - writeXRPCError(w, "ProfileNotFound", "user not found", http.StatusNotFound) + // Only an actually-absent handle is a 404. Reporting a DNS outage or + // a broken PLC directory as "user not found" tells the caller their + // handle is wrong and hides the outage from anyone watching 5xx. + var invalidHandle *users.InvalidHandleError + switch { + case errors.Is(err, users.ErrUserNotFound): + writeXRPCError(w, "ProfileNotFound", "user not found", http.StatusNotFound) + case errors.As(err, &invalidHandle): + writeXRPCError(w, "InvalidRequest", "actor is not a valid handle or DID", http.StatusBadRequest) + default: + log.Printf("Failed to resolve handle %s: %v", actor, err) + writeXRPCError(w, "InternalError", "failed to resolve handle", http.StatusInternalServerError) + } return } did = resolvedDID @@ -174,16 +187,12 @@ func (h *UserHandler) GetProfile(w http.ResponseWriter, r *http.Request) { } } -// writeXRPCError writes a standardized XRPC error response +// writeXRPCError writes a standardized XRPC error response. +// +// Argument order differs from xrpc.WriteError to match this file's existing +// call sites. func writeXRPCError(w http.ResponseWriter, errorName, message string, statusCode int) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(statusCode) - if err := json.NewEncoder(w).Encode(map[string]interface{}{ - "error": errorName, - "message": message, - }); err != nil { - log.Printf("Failed to encode error response: %v", err) - } + xrpc.WriteError(w, statusCode, errorName, message) } // Signup handles social.coves.actor.signup diff --git a/internal/api/xrpc/errors.go b/internal/api/xrpc/errors.go index 9cb1371..fa81dc6 100644 --- a/internal/api/xrpc/errors.go +++ b/internal/api/xrpc/errors.go @@ -2,7 +2,7 @@ package xrpc import ( "encoding/json" - "log" + "log/slog" "net/http" ) @@ -12,14 +12,41 @@ type Error struct { Message string `json:"message"` } -// WriteError writes an XRPC error response with the given status code +// WriteError writes an XRPC error response with the given status code. +// +// The body is marshalled before any header is sent, so an encoding failure +// cannot leave the client with a committed status and a truncated body. That +// should be unreachable for two plain strings, but the fallback keeps the +// response well-formed if it ever happens. func WriteError(w http.ResponseWriter, statusCode int, errorType, message string) { + // A malformed mapping is a bug in the rule table, but net/http panics on an + // out-of-range status and an empty code would leave the client with no + // contract to switch on. Answer a plain 500 and record the bug instead of + // taking down the request goroutine. + if statusCode < 100 || statusCode > 599 || errorType == "" { + slog.Error("invalid XRPC error mapping; substituting 500", + "status", statusCode, "errorType", errorType) + statusCode, errorType, message = http.StatusInternalServerError, + "InternalServerError", "An internal error occurred" + } + + body, err := json.Marshal(Error{Error: errorType, Message: message}) + if err != nil { + slog.Error("failed to marshal XRPC error response", + "error", err, "errorType", errorType) + w.Header().Set("Content-Type", "text/plain; charset=utf-8") + w.WriteHeader(statusCode) + if _, err := w.Write([]byte(message)); err != nil { + slog.Warn("failed to write XRPC error fallback response", "error", err) + } + return + } + w.Header().Set("Content-Type", "application/json") w.WriteHeader(statusCode) - if err := json.NewEncoder(w).Encode(Error{ - Error: errorType, - Message: message, - }); err != nil { - log.Printf("Failed to encode XRPC error response: %v", err) + if _, err := w.Write(body); err != nil { + // Status and body are already committed; the client most likely + // disconnected. Nothing to do but record it. + slog.Warn("failed to write XRPC error response", "error", err) } } diff --git a/internal/api/xrpc/mapper.go b/internal/api/xrpc/mapper.go new file mode 100644 index 0000000..bef653f --- /dev/null +++ b/internal/api/xrpc/mapper.go @@ -0,0 +1,278 @@ +package xrpc + +import ( + "context" + "errors" + "log/slog" + "net/http" + + "Coves/internal/atproto/pds" + coreerrors "Coves/internal/core/errors" +) + +// Mapping is the XRPC error response a Go error resolves to. +// +// Code is the machine-readable contract clients switch on, so it must stay +// stable; Message is human-readable and may be reworded freely. Message must +// never carry internal detail — see the *Detail rule constructors for the +// narrow cases where an error's own text is safe to surface. +type Mapping struct { + Status int + Code string + Message string +} + +// Rule reports the Mapping for err, or false when it does not apply. +type Rule func(err error) (Mapping, bool) + +// Sentinel matches err against target with errors.Is. +// +// errors.Is rather than ==: these sentinels travel up through service layers +// that add context with %w, and an == comparison silently stops matching the +// moment anyone wraps them — degrading a considered 4xx into a 500. +func Sentinel(target error, status int, code, message string) Rule { + return func(err error) (Mapping, bool) { + if !errors.Is(err, target) { + return Mapping{}, false + } + return Mapping{Status: status, Code: code, Message: message}, true + } +} + +// SentinelDetail is Sentinel with the error's own text as the message. Same +// caution as MatchDetail: only for errors written for the client to read. +func SentinelDetail(target error, status int, code string) Rule { + return func(err error) (Mapping, bool) { + if !errors.Is(err, target) { + return Mapping{}, false + } + return Mapping{Status: status, Code: code, Message: err.Error()}, true + } +} + +// Match answers with a fixed message when pred accepts err. Use it for the +// predicate helpers domain packages already expose (IsValidationError and +// friends). +func Match(pred func(error) bool, status int, code, message string) Rule { + return func(err error) (Mapping, bool) { + if !pred(err) { + return Mapping{}, false + } + return Mapping{Status: status, Code: code, Message: message}, true + } +} + +// MatchDetail is Match with the error's own text as the message. +// +// Only for errors whose text is written for the client: validation failures +// naming a bad field, say. Anything that could contain a query, a DID we did +// not receive from the caller, or a driver message belongs in Match with a +// fixed string. +func MatchDetail(pred func(error) bool, status int, code string) Rule { + return func(err error) (Mapping, bool) { + if !pred(err) { + return Mapping{}, false + } + return Mapping{Status: status, Code: code, Message: err.Error()}, true + } +} + +// As matches the first error of type T in err's chain and derives the message +// from that error alone. +// +// Preferred over MatchDetail whenever a typed error is available: it reads the +// message off the typed error itself, so wrapper context added upstream +// ("failed to create comment: ...") cannot leak into the client's message. +// +// T must be the concrete type the domain returns, pointer included +// (*posts.ValidationError). Two ways to get it wrong: instantiating with a +// value type whose Error method has a pointer receiver won't compile, and +// As[error] compiles but matches every error in existence, silently becoming a +// catch-all. +func As[T error](status int, code string, message func(T) string) Rule { + return func(err error) (Mapping, bool) { + var target T + if !errors.As(err, &target) { + return Mapping{}, false + } + return Mapping{Status: status, Code: code, Message: message(target)}, true + } +} + +// defaultReauth is the answer when the user's session is dead. Overridable per +// mapper via Mapper.WithReauth, because a few endpoints ship different codes to +// clients already. +var defaultReauth = Mapping{ + Status: http.StatusUnauthorized, + Code: "AuthRequired", + Message: "Authentication required or session expired", +} + +// internalError is the last resort. Its message is deliberately uniform and +// content-free: the cause is logged, never sent. +var internalError = Mapping{ + Status: http.StatusInternalServerError, + Code: "InternalServerError", + Message: "An internal error occurred", +} + +// sharedRules are consulted after a mapper's own rules, in this order. They +// cover the errors that reach every handler identically. +var sharedRules = []Rule{ + // Typed domain errors, read off the typed value so wrapper context cannot + // leak. One entry each, rather than one per domain package, because the + // domains alias these types — see internal/core/errors. + As[*coreerrors.ValidationError](http.StatusBadRequest, "InvalidRequest", + func(e *coreerrors.ValidationError) string { return e.Error() }), + As[*coreerrors.NotFoundError](http.StatusNotFound, "NotFound", + func(e *coreerrors.NotFoundError) string { return e.Error() }), + As[*coreerrors.ConflictError](http.StatusConflict, "AlreadyExists", + func(e *coreerrors.ConflictError) string { return e.Error() }), + + // PDS failures other than a dead session, which is handled ahead of the + // domain rules. 403 is a permissions problem — a missing OAuth scope, say — + // not an expired session, so it must not trigger a client sign-out; signing + // in again is nonetheless what re-grants the scope. + Sentinel(pds.ErrForbidden, http.StatusForbidden, "PermissionDenied", + "Your session does not have permission for this action. Sign out and back in to grant it."), + Sentinel(pds.ErrBadRequest, http.StatusBadRequest, "InvalidRequest", + "Invalid request to PDS"), + Sentinel(pds.ErrNotFound, http.StatusNotFound, "NotFound", + "Record not found on PDS"), + Sentinel(pds.ErrConflict, http.StatusConflict, "Conflict", + "Record was modified by another operation"), + Sentinel(pds.ErrPayloadTooLarge, http.StatusRequestEntityTooLarge, "PayloadTooLarge", + "Request payload exceeds size limit"), + Sentinel(pds.ErrRateLimited, http.StatusTooManyRequests, "RateLimitExceeded", + "Too many requests, please try again later"), + + // Request lifecycle. A cancellation is the client's own doing, so it is a + // 4xx; a deadline we blew is ours to report as a gateway timeout. + Sentinel(context.DeadlineExceeded, http.StatusGatewayTimeout, "Timeout", + "Request timed out"), + Sentinel(context.Canceled, http.StatusBadRequest, "RequestCanceled", + "Request was canceled"), +} + +// Mapper resolves domain errors to XRPC error responses for one handler +// package. +// +// Every handler package used to carry its own near-identical copy of this +// switch. Because each copy mapped a different subset, an error a package +// forgot fell through to a 500 — most damagingly a dead OAuth session on +// post, comment, and vote, which left clients with no signal to re-authenticate +// and no way out but a manual sign-out. A mapper owns only the rules unique to +// its domain and inherits the rest. +type Mapper struct { + domain string + rules []Rule + reauth Mapping + internal Mapping +} + +// NewMapper builds a mapper for a handler package. domain names it for logs. +// rules are the domain-specific mappings, tried in order. +func NewMapper(domain string, rules ...Rule) *Mapper { + return &Mapper{domain: domain, rules: rules, reauth: defaultReauth, internal: internalError} +} + +// WithFallback overrides the answer for an error no rule matched, returning a +// derived mapper. +// +// Two uses. An operation that already ships a more specific 500 code to +// clients — a blob upload that failed for reasons we could not classify. And a +// terminal step where every possible failure has the same meaning: if building +// a PDS client for a user fails at all, there is no usable session, whatever +// the cause, so 401 is the honest answer rather than a guess. +// +// Prefer adding a rule for the cause over widening the fallback. Whatever the +// status, an unmatched error is still logged. +func (m *Mapper) WithFallback(status int, code, message string) *Mapper { + derived := *m + derived.internal = Mapping{Status: status, Code: code, Message: message} + return &derived +} + +// WithReauth overrides the code and message used when the session is dead, +// returning a derived mapper. The status stays 401 — that is the part clients +// act on. +func (m *Mapper) WithReauth(code, message string) *Mapper { + derived := *m + derived.reauth = Mapping{Status: http.StatusUnauthorized, Code: code, Message: message} + return &derived +} + +// With returns a derived mapper whose extra rules are tried before the domain +// rules, for a single operation that needs a more specific answer than its +// package's default — distinguishing an oversized avatar from an oversized +// banner, say. +// +// The extra rules cannot displace the re-authentication check, which stays +// ahead of everything. +func (m *Mapper) With(rules ...Rule) *Mapper { + derived := *m + // Fresh backing array: appending onto m.rules would let two derived + // mappers overwrite each other's rules. + derived.rules = make([]Rule, 0, len(rules)+len(m.rules)) + derived.rules = append(derived.rules, rules...) + derived.rules = append(derived.rules, m.rules...) + return &derived +} + +// Resolve returns the response for err and reports whether any rule matched. +// +// A false result means no rule matched and the Mapping is this mapper's +// fallback — the generic 500 unless WithFallback changed it — so it says +// nothing about the cause and the caller must log err. Rules are tried in this +// order: +// +// 1. re-authentication required +// 2. the mapper's own rules, in the order given +// 3. shared rules: typed domain errors, other PDS failures, request lifecycle +// +// Re-authentication comes first on purpose. Services translate a PDS auth +// failure into their own sentinel — votes.ErrNotAuthorized, say — which maps to +// 403; were that consulted first, an expired session would answer 403 and the +// client would never learn to sign in again. Nothing else needs to jump the +// queue, so a domain rule still beats every shared rule. +func (m *Mapper) Resolve(err error) (Mapping, bool) { + if err == nil { + return m.internal, false + } + + if pds.IsReauthRequired(err) { + return m.reauth, true + } + + for _, rule := range m.rules { + if mapping, ok := rule(err); ok { + return mapping, true + } + } + + for _, rule := range sharedRules { + if mapping, ok := rule(err); ok { + return mapping, true + } + } + + return m.internal, false +} + +// Write resolves err and writes the response, logging anything unmapped. +func (m *Mapper) Write(w http.ResponseWriter, err error) { + if err == nil { + // A handler reached its error path without an error. Answering 500 is + // wrong but at least terminates the request; a silent return would + // leave the client with an empty 200. + slog.Error("error mapper invoked with nil error", "domain", m.domain) + WriteError(w, m.internal.Status, m.internal.Code, m.internal.Message) + return + } + + mapping, matched := m.Resolve(err) + if !matched { + slog.Error("unmapped handler error", "domain", m.domain, "error", err) + } + WriteError(w, mapping.Status, mapping.Code, mapping.Message) +} diff --git a/internal/api/xrpc/mapper_test.go b/internal/api/xrpc/mapper_test.go new file mode 100644 index 0000000..3391cb9 --- /dev/null +++ b/internal/api/xrpc/mapper_test.go @@ -0,0 +1,319 @@ +package xrpc_test + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "net/http/httptest" + "testing" + + "Coves/internal/api/xrpc" + "Coves/internal/atproto/pds" + coreerrors "Coves/internal/core/errors" +) + +// Stand-ins for a domain package's sentinels. +var ( + errDomainNotAuthorized = errors.New("not authorized") + errDomainNotFound = errors.New("thing not found") +) + +func newTestMapper() *xrpc.Mapper { + return xrpc.NewMapper("test", + xrpc.Sentinel(errDomainNotAuthorized, http.StatusForbidden, + "NotAuthorized", "You may not do that"), + xrpc.Sentinel(errDomainNotFound, http.StatusNotFound, + "ThingNotFound", "Thing not found"), + ) +} + +func TestResolve(t *testing.T) { + tests := []struct { + name string + err error + wantStatus int + wantCode string + wantMatch bool + }{ + { + name: "domain sentinel", + err: errDomainNotAuthorized, + wantStatus: http.StatusForbidden, + wantCode: "NotAuthorized", + wantMatch: true, + }, + { + // The regression this package exists to prevent: a wrapped sentinel + // must still match. An == comparison would drop to 500 here. + name: "domain sentinel wrapped with %w", + err: fmt.Errorf("service call failed: %w", errDomainNotFound), + wantStatus: http.StatusNotFound, + wantCode: "ThingNotFound", + wantMatch: true, + }, + { + name: "pds 401 with no domain sentinel", + err: fmt.Errorf("CreateRecord: %w: token expired", pds.ErrUnauthorized), + wantStatus: http.StatusUnauthorized, + wantCode: "AuthRequired", + wantMatch: true, + }, + { + // A session that could not be resumed never reaches the PDS to earn + // a 401, but has the same remedy. + name: "session expired", + err: fmt.Errorf("failed to resume: %w: gone", pds.ErrSessionExpired), + wantStatus: http.StatusUnauthorized, + wantCode: "AuthRequired", + wantMatch: true, + }, + { + // 403 is a scope problem, not a dead session; it must not answer 401 + // or clients would sign the user out over a permissions gap. + name: "pds 403 stays 403", + err: fmt.Errorf("CreateRecord: %w: no scope", pds.ErrForbidden), + wantStatus: http.StatusForbidden, + wantCode: "PermissionDenied", + wantMatch: true, + }, + { + name: "pds 429 inherited", + err: fmt.Errorf("CreateRecord: %w", pds.ErrRateLimited), + wantStatus: http.StatusTooManyRequests, + wantCode: "RateLimitExceeded", + wantMatch: true, + }, + { + name: "pds 413 inherited", + err: fmt.Errorf("UploadBlob: %w", pds.ErrPayloadTooLarge), + wantStatus: http.StatusRequestEntityTooLarge, + wantCode: "PayloadTooLarge", + wantMatch: true, + }, + { + name: "shared typed validation error", + err: coreerrors.NewValidationError("handle", "is required"), + wantStatus: http.StatusBadRequest, + wantCode: "InvalidRequest", + wantMatch: true, + }, + { + name: "shared typed not found", + err: coreerrors.NewNotFoundError("post", "at://x"), + wantStatus: http.StatusNotFound, + wantCode: "NotFound", + wantMatch: true, + }, + { + name: "shared typed conflict", + err: coreerrors.NewConflictError("community", "handle", "!go"), + wantStatus: http.StatusConflict, + wantCode: "AlreadyExists", + wantMatch: true, + }, + { + name: "deadline exceeded", + err: fmt.Errorf("query: %w", context.DeadlineExceeded), + wantStatus: http.StatusGatewayTimeout, + wantCode: "Timeout", + wantMatch: true, + }, + { + name: "canceled", + err: fmt.Errorf("query: %w", context.Canceled), + wantStatus: http.StatusBadRequest, + wantCode: "RequestCanceled", + wantMatch: true, + }, + { + name: "unmapped falls through and reports no match", + err: errors.New("something nobody anticipated"), + wantStatus: http.StatusInternalServerError, + wantCode: "InternalServerError", + wantMatch: false, + }, + } + + mapper := newTestMapper() + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, matched := mapper.Resolve(tt.err) + if got.Status != tt.wantStatus { + t.Errorf("status = %d, want %d", got.Status, tt.wantStatus) + } + if got.Code != tt.wantCode { + t.Errorf("code = %q, want %q", got.Code, tt.wantCode) + } + if matched != tt.wantMatch { + t.Errorf("matched = %v, want %v", matched, tt.wantMatch) + } + }) + } +} + +// TestReauthBeatsDomainSentinel is the fix for the bug that motivated this +// package. Services translate a PDS auth failure into their own sentinel, which +// maps to 403. If that were consulted first, a user whose session expired would +// get 403 forever with no signal to sign in again. +func TestReauthBeatsDomainSentinel(t *testing.T) { + mapper := newTestMapper() + + // What a service now returns: its own sentinel, with the cause preserved. + expired := fmt.Errorf("%w: %w", + errDomainNotAuthorized, + fmt.Errorf("CreateRecord: %w: token expired", pds.ErrUnauthorized)) + + // Both are matchable... + if !errors.Is(expired, errDomainNotAuthorized) { + t.Fatal("domain sentinel should still match") + } + if !errors.Is(expired, pds.ErrUnauthorized) { + t.Fatal("pds cause should still match") + } + + // ...and the dead session wins. + got, matched := mapper.Resolve(expired) + if !matched { + t.Fatal("expected a match") + } + if got.Status != http.StatusUnauthorized { + t.Errorf("status = %d, want 401 so the client re-authenticates", got.Status) + } + if got.Code != "AuthRequired" { + t.Errorf("code = %q, want AuthRequired", got.Code) + } + + // A 403 layered the same way keeps the domain's own answer, since both mean + // "not allowed" and neither is fixed by signing in again. + forbidden := fmt.Errorf("%w: %w", + errDomainNotAuthorized, + fmt.Errorf("CreateRecord: %w", pds.ErrForbidden)) + got, _ = mapper.Resolve(forbidden) + if got.Status != http.StatusForbidden { + t.Errorf("status = %d, want 403", got.Status) + } + if got.Code != "NotAuthorized" { + t.Errorf("code = %q, want the domain code NotAuthorized", got.Code) + } +} + +func TestRuleOrderIsFirstMatchWins(t *testing.T) { + specific := errors.New("specific") + mapper := xrpc.NewMapper("test", + xrpc.Sentinel(specific, http.StatusNotFound, "Specific", "specific"), + xrpc.Match(func(error) bool { return true }, http.StatusBadRequest, "Broad", "broad"), + ) + + if got, _ := mapper.Resolve(specific); got.Code != "Specific" { + t.Errorf("code = %q, want Specific: an earlier rule must win", got.Code) + } + if got, _ := mapper.Resolve(errors.New("other")); got.Code != "Broad" { + t.Errorf("code = %q, want Broad", got.Code) + } +} + +func TestWithDerivesWithoutMutatingParent(t *testing.T) { + parent := newTestMapper() + derived := parent.With( + xrpc.Sentinel(errDomainNotFound, http.StatusNotFound, "Narrower", "narrower"), + ) + + if got, _ := derived.Resolve(errDomainNotFound); got.Code != "Narrower" { + t.Errorf("derived code = %q, want Narrower", got.Code) + } + if got, _ := parent.Resolve(errDomainNotFound); got.Code != "ThingNotFound" { + t.Errorf("parent code = %q, want ThingNotFound: With must not mutate the parent", got.Code) + } + + // Two siblings must not share a backing array. + siblingA := parent.With(xrpc.Sentinel(errDomainNotFound, http.StatusNotFound, "A", "a")) + siblingB := parent.With(xrpc.Sentinel(errDomainNotFound, http.StatusNotFound, "B", "b")) + if got, _ := siblingA.Resolve(errDomainNotFound); got.Code != "A" { + t.Errorf("siblingA code = %q, want A", got.Code) + } + if got, _ := siblingB.Resolve(errDomainNotFound); got.Code != "B" { + t.Errorf("siblingB code = %q, want B", got.Code) + } +} + +// A derived mapper's extra rules must not be able to displace the +// re-authentication check. +func TestWithCannotDisplaceReauth(t *testing.T) { + mapper := newTestMapper().With( + xrpc.Match(func(error) bool { return true }, http.StatusTeapot, "Greedy", "greedy"), + ) + + got, _ := mapper.Resolve(fmt.Errorf("x: %w", pds.ErrUnauthorized)) + if got.Status != http.StatusUnauthorized { + t.Errorf("status = %d, want 401 even behind a catch-all rule", got.Status) + } +} + +func TestWithReauthAndWithFallback(t *testing.T) { + mapper := newTestMapper(). + WithReauth("AuthExpired", "Please re-authenticate."). + WithFallback(http.StatusUnauthorized, "SessionError", "Sign in again.") + + got, _ := mapper.Resolve(fmt.Errorf("x: %w", pds.ErrUnauthorized)) + if got.Code != "AuthExpired" || got.Status != http.StatusUnauthorized { + t.Errorf("reauth = %d/%q, want 401/AuthExpired", got.Status, got.Code) + } + + got, matched := mapper.Resolve(errors.New("unclassifiable")) + if got.Code != "SessionError" || got.Status != http.StatusUnauthorized { + t.Errorf("fallback = %d/%q, want 401/SessionError", got.Status, got.Code) + } + if matched { + t.Error("a fallback answer must still report matched=false so the caller logs it") + } +} + +func TestAsReadsMessageOffTypedErrorNotTheChain(t *testing.T) { + mapper := xrpc.NewMapper("test") + wrapped := fmt.Errorf("internal detail nobody should see: %w", + coreerrors.NewValidationError("handle", "is required")) + + got, _ := mapper.Resolve(wrapped) + if got.Message != "handle: is required" { + t.Errorf("message = %q, want just the typed error's text", got.Message) + } +} + +func TestWriteProducesXRPCJSON(t *testing.T) { + rec := httptest.NewRecorder() + newTestMapper().Write(rec, errDomainNotAuthorized) + + if rec.Code != http.StatusForbidden { + t.Errorf("status = %d, want 403", rec.Code) + } + if ct := rec.Header().Get("Content-Type"); ct != "application/json" { + t.Errorf("content-type = %q, want application/json", ct) + } + + var body xrpc.Error + if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { + t.Fatalf("body is not valid JSON: %v", err) + } + if body.Error != "NotAuthorized" { + t.Errorf("error = %q, want NotAuthorized", body.Error) + } + if body.Message != "You may not do that" { + t.Errorf("message = %q", body.Message) + } +} + +// A nil error means the handler took its error path without an error. Answering +// 500 is wrong, but a silent return would leave the client with an empty 200. +func TestWriteNilErrorStillTerminatesTheRequest(t *testing.T) { + rec := httptest.NewRecorder() + newTestMapper().Write(rec, nil) + + if rec.Code != http.StatusInternalServerError { + t.Errorf("status = %d, want 500", rec.Code) + } + if rec.Body.Len() == 0 { + t.Error("expected a body") + } +} diff --git a/internal/atproto/identity/postgres_cache.go b/internal/atproto/identity/postgres_cache.go index e54255d..3be8549 100644 --- a/internal/atproto/identity/postgres_cache.go +++ b/internal/atproto/identity/postgres_cache.go @@ -3,6 +3,7 @@ package identity import ( "context" "database/sql" + "errors" "fmt" "log" "strings" @@ -46,7 +47,7 @@ func (r *postgresCache) Get(ctx context.Context, identifier string) (*Identity, &expiresAt, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, &ErrCacheMiss{Identifier: identifier} } if err != nil { diff --git a/internal/atproto/jetstream/comment_consumer.go b/internal/atproto/jetstream/comment_consumer.go index 1247c94..1844a0f 100644 --- a/internal/atproto/jetstream/comment_consumer.go +++ b/internal/atproto/jetstream/comment_consumer.go @@ -7,6 +7,7 @@ import ( "context" "database/sql" "encoding/json" + "errors" "fmt" "log" "strings" @@ -83,7 +84,7 @@ func (c *CommentEventConsumer) bridgeStatsAllowedForRepo(ctx context.Context, re var pdsURL string err := c.db.QueryRowContext(ctx, `SELECT pds_url FROM users WHERE did = $1`, repoDID).Scan(&pdsURL) if err != nil { - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { log.Printf("debug: ignoring bridgedStats from repo %s (user not indexed; provenance unverifiable)", repoDID) } else { log.Printf("Warning: bridgedStats provenance check failed for repo %s: %v", repoDID, err) @@ -244,7 +245,7 @@ func (c *CommentEventConsumer) updateComment(ctx context.Context, repoDID string `SELECT root_uri, root_cid, parent_uri, parent_cid, deleted_at, bridged_stats_as_of, indexed_at FROM comments WHERE uri = $1`, uri, ).Scan(&storedRootURI, &storedRootCID, &storedParentURI, &storedParentCID, &storedDeletedAt, &storedAsOf, &storedIndexedAt) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { // Comment not indexed yet: its CREATE event will index it when it arrives. log.Printf("Update event for non-indexed comment: %s (will be indexed on CREATE)", uri) return nil @@ -684,7 +685,7 @@ func (c *CommentEventConsumer) indexCommentAndUpdateCounts(ctx context.Context, return fmt.Errorf("failed to resurrect comment: %w", err) } - } else if checkErr == sql.ErrNoRows { + } else if errors.Is(checkErr, sql.ErrNoRows) { // Comment doesn't exist - insert new comment // Use ON CONFLICT DO NOTHING to handle race conditions gracefully // (e.g., duplicate Jetstream events from reconnections/retries) @@ -714,7 +715,7 @@ func (c *CommentEventConsumer) indexCommentAndUpdateCounts(ctx context.Context, comment.CreatedAt, comment.IndexedAt, comment.BridgedUpvoteCount, comment.BridgedDownvoteCount, comment.BridgedStatsAsOf, comment.Score, ).Scan(&commentID) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { // ON CONFLICT triggered - comment was inserted by concurrent process // This is an idempotent replay, skip gracefully log.Printf("Comment already indexed (concurrent insert): %s", comment.URI) diff --git a/internal/atproto/jetstream/post_consumer.go b/internal/atproto/jetstream/post_consumer.go index 603cb5e..7ad9651 100644 --- a/internal/atproto/jetstream/post_consumer.go +++ b/internal/atproto/jetstream/post_consumer.go @@ -339,7 +339,7 @@ func (c *PostEventConsumer) updatePost(ctx context.Context, repoDID string, comm `SELECT id, community_did, author_did, deleted_at, bridged_stats_as_of, indexed_at FROM posts WHERE uri = $1`, uri, ).Scan(&storedID, &storedCommunityDID, &storedAuthorDID, &storedDeletedAt, &storedAsOf, &storedIndexedAt) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { // Not indexed yet (out-of-order delivery). Jetstream will replay CREATE; skip. log.Printf("Update event for non-indexed post: %s (will be indexed on CREATE)", uri) return nil @@ -626,7 +626,7 @@ func (c *PostEventConsumer) indexPostAndReconcileCounts(ctx context.Context, pos ).Scan(&postID) // If no rows returned, post already exists (idempotent - OK for Jetstream replays) - if insertErr == sql.ErrNoRows { + if errors.Is(insertErr, sql.ErrNoRows) { // KNOWN LIMITATION (accepted): a genuine RE-CREATE of the same rkey while // the row is still ACTIVE also lands here and is treated as an idempotent // duplicate, dropping the new content. Reaching that state requires the diff --git a/internal/atproto/jetstream/rev_gate.go b/internal/atproto/jetstream/rev_gate.go index 97c50b4..10022f6 100644 --- a/internal/atproto/jetstream/rev_gate.go +++ b/internal/atproto/jetstream/rev_gate.go @@ -3,6 +3,7 @@ package jetstream import ( "context" "database/sql" + "errors" "fmt" "log" ) @@ -90,7 +91,7 @@ func recordRevIsStale(ctx context.Context, q revGateQuerier, uri, rev string) (b err := q.QueryRowContext(ctx, `SELECT rev >= $2 FROM jetstream_record_revs WHERE record_uri = $1`, uri, rev, ).Scan(&stale) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return false, nil } if err != nil { diff --git a/internal/atproto/jetstream/vote_consumer.go b/internal/atproto/jetstream/vote_consumer.go index eaa666b..abc8a56 100644 --- a/internal/atproto/jetstream/vote_consumer.go +++ b/internal/atproto/jetstream/vote_consumer.go @@ -6,6 +6,7 @@ import ( "Coves/internal/core/votes" "context" "database/sql" + "errors" "fmt" "log" "strings" @@ -160,7 +161,7 @@ func (c *VoteEventConsumer) deleteVote(ctx context.Context, repoDID string, comm err = tx.QueryRowContext(ctx, `SELECT direction, subject_uri, deleted_at FROM votes WHERE uri = $1`, uri, ).Scan(&direction, &subjectURI, &deletedAt) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { // Idempotent: vote never indexed. Commit the gate advance anyway — it // is the tombstone that rejects a stale cross-feed copy of the CREATE // arriving later for a record that no longer exists on the PDS. @@ -392,7 +393,7 @@ func (c *VoteEventConsumer) indexVoteAndUpdateCounts(ctx context.Context, vote * ).Scan(&voteID) // If no rows returned, vote already exists (idempotent - OK for Jetstream replays) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { // KNOWN LIMITATION (accepted): a genuine RE-CREATE of the same rkey while // the row is still ACTIVE also lands here and is treated as an idempotent // duplicate. Reaching that state requires the exact sequence: create A diff --git a/internal/atproto/oauth/store.go b/internal/atproto/oauth/store.go index 4e050ae..5ee2daf 100644 --- a/internal/atproto/oauth/store.go +++ b/internal/atproto/oauth/store.go @@ -70,7 +70,7 @@ func (s *PostgresOAuthStore) GetSession(ctx context.Context, did syntax.DID, ses &dpopPrivateKeyMultibase, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, ErrSessionNotFound } if err != nil { @@ -284,7 +284,7 @@ func (s *PostgresOAuthStore) GetAuthRequestInfo(ctx context.Context, state strin &createdAt, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, ErrAuthRequestNotFound } if err != nil { @@ -624,7 +624,7 @@ func (s *PostgresOAuthStore) GetMobileOAuthData(ctx context.Context, state strin var csrfToken, redirectURI sql.NullString err := s.db.QueryRowContext(ctx, query, state).Scan(&csrfToken, &redirectURI) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, ErrAuthRequestNotFound } if err != nil { diff --git a/internal/atproto/pds/errors.go b/internal/atproto/pds/errors.go index 91c5962..d891736 100644 --- a/internal/atproto/pds/errors.go +++ b/internal/atproto/pds/errors.go @@ -26,12 +26,33 @@ var ( // ErrPayloadTooLarge indicates the request payload exceeds PDS limits (HTTP 413). ErrPayloadTooLarge = errors.New("payload too large") + + // ErrSessionExpired indicates a stored OAuth session could not be resumed: + // the refresh token expired, the session was revoked on the PDS, or the + // DPoP key no longer matches. Unlike ErrUnauthorized this is detected + // locally, before any request reaches the PDS, so it carries no HTTP + // status — but it has the same remedy, and a client that is not told to + // re-authenticate will retry forever. + ErrSessionExpired = errors.New("oauth session expired") ) // IsAuthError returns true if the error is an authentication/authorization error. -// This is a convenience function for checking if re-authentication might help. +// +// It deliberately spans both 401 and 403, so it must NOT be used to pick an +// HTTP status: 401 means "sign in again" while 403 means "your session lacks +// the scope for this", and collapsing them leaves clients unable to tell a +// dead session from a permissions gap. Use IsReauthRequired for that decision, +// or match the individual sentinels. func IsAuthError(err error) bool { - return errors.Is(err, ErrUnauthorized) || errors.Is(err, ErrForbidden) + return errors.Is(err, ErrUnauthorized) || errors.Is(err, ErrForbidden) || errors.Is(err, ErrSessionExpired) +} + +// IsReauthRequired reports whether the user's session is no longer usable and +// the client must start a new sign-in. This is the check that should drive a +// 401 response. ErrForbidden is excluded: re-authenticating with the same +// scopes would fail the same way. +func IsReauthRequired(err error) bool { + return errors.Is(err, ErrUnauthorized) || errors.Is(err, ErrSessionExpired) } // IsConflictError returns true if the error indicates a conflict (e.g., duplicate record). diff --git a/internal/atproto/pds/factory.go b/internal/atproto/pds/factory.go index d725ba8..660dbac 100644 --- a/internal/atproto/pds/factory.go +++ b/internal/atproto/pds/factory.go @@ -2,9 +2,12 @@ package pds import ( "context" + "errors" "fmt" "net/http" + covesoauth "Coves/internal/atproto/oauth" + "github.com/bluesky-social/indigo/atproto/atclient" "github.com/bluesky-social/indigo/atproto/auth/oauth" "github.com/bluesky-social/indigo/atproto/syntax" @@ -32,9 +35,8 @@ func NewFromOAuthSession(ctx context.Context, oauthClient *oauth.ClientApp, sess // - DPoP key mismatch → Session data corrupted, re-authenticate sess, err := oauthClient.ResumeSession(ctx, sessionData.AccountDID, sessionData.SessionID) if err != nil { - // Include DID and session context for debugging - return nil, fmt.Errorf("failed to resume OAuth session for DID=%s, sessionID=%s: %w", - sessionData.AccountDID.String(), sessionData.SessionID, err) + return nil, classifyResumeFailure(err, + sessionData.AccountDID.String(), sessionData.SessionID) } // APIClient() returns an *atclient.APIClient configured with DPoP auth @@ -47,6 +49,31 @@ func NewFromOAuthSession(ctx context.Context, oauthClient *oauth.ClientApp, sess }, nil } +// classifyResumeFailure decides whether a failed session resume means the user +// must sign in again. +// +// Tag ONLY a session that is genuinely gone. ResumeSession is a session-store +// read and nothing more, so its failures split two ways: the row is absent or +// past its expiry (terminal — signing in again is the fix), or the store itself +// failed (a database outage, an exhausted pool, a cancelled request). +// +// The distinction has to be made here because the API boundary checks +// re-authentication ahead of every other rule, so anything tagged expired +// answers 401. Tagging the whole class would turn a few seconds of database +// trouble into a sign-out for every user with a request in flight, and would +// hide the outage from 5xx alerting at the same time. +// +// Either way the cause stays wrapped, so it reaches the logs and — for a +// cancelled or timed-out request — still matches the boundary's lifecycle rules. +func classifyResumeFailure(err error, did, sessionID string) error { + if errors.Is(err, covesoauth.ErrSessionNotFound) { + return fmt.Errorf("failed to resume OAuth session for DID=%s, sessionID=%s: %w: %w", + did, sessionID, ErrSessionExpired, err) + } + return fmt.Errorf("failed to resume OAuth session for DID=%s, sessionID=%s: %w", + did, sessionID, err) +} + // NewFromPasswordAuth creates a PDS client using password authentication. // This uses Bearer token authentication from com.atproto.server.createSession. // diff --git a/internal/atproto/pds/factory_session_test.go b/internal/atproto/pds/factory_session_test.go new file mode 100644 index 0000000..12e9903 --- /dev/null +++ b/internal/atproto/pds/factory_session_test.go @@ -0,0 +1,136 @@ +package pds + +import ( + "context" + "errors" + "fmt" + "testing" + + covesoauth "Coves/internal/atproto/oauth" +) + +// TestIsReauthRequired pins which conditions mean "the user must sign in again". +// +// This predicate is the first thing the API error mapper consults, and it is the +// only thing standing between a real expired session and a spurious sign-out. +// The ErrForbidden row is the important one: a missing OAuth scope is not a dead +// session, and answering 401 for it would sign users out over a permissions gap +// that signing in again does not obviously fix. +func TestIsReauthRequired(t *testing.T) { + tests := []struct { + name string + err error + want bool + }{ + {"nil", nil, false}, + {"unauthorized", ErrUnauthorized, true}, + {"session expired", ErrSessionExpired, true}, + {"forbidden is NOT reauth", ErrForbidden, false}, + {"not found", ErrNotFound, false}, + {"bad request", ErrBadRequest, false}, + {"rate limited", ErrRateLimited, false}, + {"unrelated", errors.New("boom"), false}, + {"wrapped unauthorized", fmt.Errorf("CreateRecord: %w: expired", ErrUnauthorized), true}, + {"wrapped session expired", fmt.Errorf("resume: %w", ErrSessionExpired), true}, + {"double-wrapped behind a domain sentinel", + fmt.Errorf("%w: %w", errors.New("not authorized"), ErrUnauthorized), true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := IsReauthRequired(tt.err); got != tt.want { + t.Errorf("IsReauthRequired(%v) = %v, want %v", tt.err, got, tt.want) + } + }) + } +} + +// TestIsAuthErrorSpansForbidden guards the distinction IsAuthError deliberately +// blurs, so nobody "simplifies" the two predicates back into one. +func TestIsAuthErrorSpansForbidden(t *testing.T) { + if !IsAuthError(ErrForbidden) { + t.Error("IsAuthError must cover 403") + } + if IsReauthRequired(ErrForbidden) { + t.Error("IsReauthRequired must NOT cover 403 — that is the whole point of having both") + } + if !IsAuthError(ErrSessionExpired) { + t.Error("IsAuthError must cover a dead session") + } +} + +// TestNewFromOAuthSessionClassifiesResumeFailures is the test whose absence made +// every "expired session returns 401" test elsewhere tautological: they all +// hand-constructed an ErrSessionExpired-tagged error rather than proving the +// factory produces one. +// +// It also pins the narrowing that matters operationally. ResumeSession is a +// session-store read, so a store outage surfaces here identically to a missing +// row — and because the mapper checks re-auth first, tagging both would answer +// 401 to every in-flight request during a database blip. +func TestNewFromOAuthSessionClassifiesResumeFailures(t *testing.T) { + tests := []struct { + name string + storeErr error + wantReauth bool + wantCauseKept bool + }{ + { + name: "missing or expired session is terminal", + storeErr: covesoauth.ErrSessionNotFound, + wantReauth: true, + wantCauseKept: true, + }, + { + name: "wrapped missing session is still terminal", + storeErr: fmt.Errorf("lookup: %w", covesoauth.ErrSessionNotFound), + wantReauth: true, + wantCauseKept: true, + }, + { + name: "store outage must NOT read as an expired session", + storeErr: errors.New("failed to get session: connection reset by peer"), + wantReauth: false, + }, + { + name: "cancelled request must NOT read as an expired session", + storeErr: context.Canceled, + wantReauth: false, + }, + { + name: "deadline must NOT read as an expired session", + storeErr: context.DeadlineExceeded, + wantReauth: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := classifyResumeFailure(tt.storeErr, "did:plc:test", "session-1") + + if got := IsReauthRequired(err); got != tt.wantReauth { + t.Errorf("IsReauthRequired = %v, want %v (err: %v)", got, tt.wantReauth, err) + } + // The cause must stay reachable either way — it is what gets logged. + if !errors.Is(err, tt.storeErr) { + t.Errorf("underlying cause is no longer matchable: %v", err) + } + if tt.wantCauseKept && !errors.Is(err, ErrSessionExpired) { + t.Error("expected ErrSessionExpired to be matchable alongside the cause") + } + }) + } +} + +// Non-terminal failures must also stay clear of the context rules being +// pre-empted: a cancelled request should still look cancelled at the boundary. +func TestResumeFailureLeavesContextErrorsIntact(t *testing.T) { + err := classifyResumeFailure(context.Canceled, "did:plc:test", "session-1") + + if !errors.Is(err, context.Canceled) { + t.Error("context.Canceled must remain matchable so the boundary can answer 400, not 401") + } + if errors.Is(err, ErrSessionExpired) { + t.Error("a cancelled request is not an expired session") + } +} diff --git a/internal/core/blueskypost/repository.go b/internal/core/blueskypost/repository.go index 63ea3fc..3662288 100644 --- a/internal/core/blueskypost/repository.go +++ b/internal/core/blueskypost/repository.go @@ -38,7 +38,7 @@ func (r *postgresBlueskyPostRepo) Get(ctx context.Context, atURI string) (*Blues var metadataJSON []byte err := r.db.QueryRowContext(ctx, query, atURI).Scan(&metadataJSON) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { // Not found or expired is a cache miss return nil, ErrCacheMiss } diff --git a/internal/core/comments/comment_service.go b/internal/core/comments/comment_service.go index c5acc5f..3ac0bc7 100644 --- a/internal/core/comments/comment_service.go +++ b/internal/core/comments/comment_service.go @@ -756,7 +756,7 @@ func (s *commentService) CreateComment(ctx context.Context, session *oauth.Clien "root", req.Reply.Root.URI, "parent", req.Reply.Parent.URI) if pds.IsAuthError(err) { - return nil, ErrNotAuthorized + return nil, fmt.Errorf("%w: %w", ErrNotAuthorized, err) } return nil, fmt.Errorf("failed to create comment: %w", err) } @@ -839,10 +839,10 @@ func (s *commentService) UpdateComment(ctx context.Context, session *oauth.Clien "uri", req.URI, "rkey", rkey) if pds.IsAuthError(err) { - return nil, ErrNotAuthorized + return nil, fmt.Errorf("%w: %w", ErrNotAuthorized, err) } if errors.Is(err, pds.ErrNotFound) { - return nil, ErrCommentNotFound + return nil, fmt.Errorf("%w: %w", ErrCommentNotFound, err) } return nil, fmt.Errorf("failed to fetch existing comment: %w", err) } @@ -891,10 +891,10 @@ func (s *commentService) UpdateComment(ctx context.Context, session *oauth.Clien "uri", req.URI, "rkey", rkey) if pds.IsAuthError(err) { - return nil, ErrNotAuthorized + return nil, fmt.Errorf("%w: %w", ErrNotAuthorized, err) } if errors.Is(err, pds.ErrConflict) { - return nil, ErrConcurrentModification + return nil, fmt.Errorf("%w: %w", ErrConcurrentModification, err) } return nil, fmt.Errorf("failed to update comment: %w", err) } @@ -951,10 +951,10 @@ func (s *commentService) DeleteComment(ctx context.Context, session *oauth.Clien "uri", req.URI, "rkey", rkey) if pds.IsAuthError(err) { - return ErrNotAuthorized + return fmt.Errorf("%w: %w", ErrNotAuthorized, err) } if errors.Is(err, pds.ErrNotFound) { - return ErrCommentNotFound + return fmt.Errorf("%w: %w", ErrCommentNotFound, err) } return fmt.Errorf("failed to verify comment: %w", err) } @@ -966,7 +966,7 @@ func (s *commentService) DeleteComment(ctx context.Context, session *oauth.Clien "uri", req.URI, "rkey", rkey) if pds.IsAuthError(err) { - return ErrNotAuthorized + return fmt.Errorf("%w: %w", ErrNotAuthorized, err) } return fmt.Errorf("failed to delete comment: %w", err) } diff --git a/internal/core/communities/service.go b/internal/core/communities/service.go index e47fc9e..ca5fdf1 100644 --- a/internal/core/communities/service.go +++ b/internal/core/communities/service.go @@ -801,7 +801,7 @@ func (s *communityService) SubscribeToCommunity(ctx context.Context, session *oa recordURI, recordCID, err := pdsClient.CreateRecord(ctx, "social.coves.community.subscription", tid.String(), subRecord) if err != nil { if pds.IsAuthError(err) { - return nil, ErrUnauthorized + return nil, fmt.Errorf("%w: %w", ErrUnauthorized, err) } return nil, fmt.Errorf("failed to create subscription on PDS: %w", err) } @@ -856,7 +856,7 @@ func (s *communityService) UnsubscribeFromCommunity(ctx context.Context, session // CRITICAL: Delete from social.coves.community.subscription (RECORD TYPE), not social.coves.community.unsubscribe if err := pdsClient.DeleteRecord(ctx, "social.coves.community.subscription", rkey); err != nil { if pds.IsAuthError(err) { - return ErrUnauthorized + return fmt.Errorf("%w: %w", ErrUnauthorized, err) } return fmt.Errorf("failed to delete subscription on PDS: %w", err) } @@ -954,7 +954,7 @@ func (s *communityService) BlockCommunity(ctx context.Context, session *oauth.Cl if err != nil { // Check for auth errors first if pds.IsAuthError(err) { - return nil, ErrUnauthorized + return nil, fmt.Errorf("%w: %w", ErrUnauthorized, err) } // Check if this is a duplicate/conflict error from PDS @@ -1027,7 +1027,7 @@ func (s *communityService) UnblockCommunity(ctx context.Context, session *oauth. // Write-forward: delete record from PDS using DPoP-authenticated client if err := pdsClient.DeleteRecord(ctx, "social.coves.community.block", rkey); err != nil { if pds.IsAuthError(err) { - return ErrUnauthorized + return fmt.Errorf("%w: %w", ErrUnauthorized, err) } return fmt.Errorf("failed to delete block on PDS: %w", err) } diff --git a/internal/core/posts/service.go b/internal/core/posts/service.go index fb10bec..877a0c9 100644 --- a/internal/core/posts/service.go +++ b/internal/core/posts/service.go @@ -922,6 +922,22 @@ func validateDIDFormat(did string) error { } } +// communityCredentialFailure reports that the *community's* stored PDS +// credentials were rejected, deliberately severing the pds sentinel from the +// chain with %v rather than %w. +// +// Posts live in the community's repo, so deletes authenticate with the +// community's service token, not the caller's OAuth session. If that token were +// allowed to surface pds.ErrUnauthorized, the API boundary would read it as the +// caller's session being dead and answer 401 — telling a user with a perfectly +// healthy session to sign in again over a server-side credential problem they +// cannot fix, and hiding a real outage from 5xx alerting. Unclassified is the +// correct answer here: it becomes a logged 500. +func communityCredentialFailure(operation, communityDID string, err error) error { + return fmt.Errorf("community PDS credentials rejected during %s for %s: %v", + operation, communityDID, err) +} + // DeletePost deletes a post from the community's PDS repository // SECURITY: Only the post author can delete their own posts // Flow: @@ -979,6 +995,9 @@ func (s *postService) DeletePost(ctx context.Context, session *oauth.ClientSessi log.Printf("[POST-DELETE] Post not found on PDS (already deleted?): %s", req.URI) return nil } + if pds.IsAuthError(err) { + return communityCredentialFailure("fetch post", community.DID, err) + } return fmt.Errorf("failed to fetch post from PDS: %w", err) } @@ -1002,6 +1021,9 @@ func (s *postService) DeletePost(ctx context.Context, session *oauth.ClientSessi log.Printf("[POST-DELETE] Post already deleted from PDS: %s", req.URI) return nil } + if pds.IsAuthError(err) { + return communityCredentialFailure("delete post", community.DID, err) + } return fmt.Errorf("failed to delete post from PDS: %w", err) } diff --git a/internal/core/unfurl/repository.go b/internal/core/unfurl/repository.go index a9e9ff6..a093c6a 100644 --- a/internal/core/unfurl/repository.go +++ b/internal/core/unfurl/repository.go @@ -4,6 +4,7 @@ import ( "context" "database/sql" "encoding/json" + "errors" "fmt" "time" ) @@ -32,7 +33,7 @@ func (r *postgresUnfurlRepo) Get(ctx context.Context, url string) (*UnfurlResult var provider string err := r.db.QueryRowContext(ctx, query, url).Scan(&metadataJSON, &thumbnailURL, &provider) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { // Not found or expired is not an error return nil, nil } diff --git a/internal/core/users/resolve_handle_test.go b/internal/core/users/resolve_handle_test.go new file mode 100644 index 0000000..aaa719a --- /dev/null +++ b/internal/core/users/resolve_handle_test.go @@ -0,0 +1,123 @@ +package users + +import ( + "context" + "errors" + "fmt" + "testing" + + "Coves/internal/atproto/identity" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" +) + +// TestResolveHandleToDID_ClassifiesResolverErrors covers the translation from +// the identity resolver's typed errors into this package's vocabulary. +// +// Without it the API boundary cannot tell "no such handle" from "resolution +// broke", and the most common failure of a public profile lookup — a handle +// that does not exist — reads as a server fault. The repo lookup deliberately +// swallows its own ErrUserNotFound and falls through to external resolution, so +// the sentinel has to be reintroduced here or it never appears at all. +func TestResolveHandleToDID_ClassifiesResolverErrors(t *testing.T) { + const handle = "ghost.example.com" + + tests := []struct { + name string + resolver error + assert func(t *testing.T, err error) + }{ + { + name: "unresolvable handle is a not-found", + resolver: &identity.ErrNotFound{Identifier: handle, Reason: "no DNS TXT record"}, + assert: func(t *testing.T, err error) { + assert.ErrorIs(t, err, ErrUserNotFound, + "an absent handle must be matchable as ErrUserNotFound so the boundary answers 404") + var notFound *identity.ErrNotFound + assert.ErrorAs(t, err, ¬Found, "the resolver's cause should stay reachable for logs") + }, + }, + { + name: "malformed handle is a client error, not a not-found", + resolver: &identity.ErrInvalidIdentifier{Identifier: "not a handle", Reason: "contains a space"}, + assert: func(t *testing.T, err error) { + var invalidHandle *InvalidHandleError + assert.ErrorAs(t, err, &invalidHandle) + assert.NotErrorIs(t, err, ErrUserNotFound, + "a malformed handle is a 400, not a 404") + }, + }, + { + name: "infrastructure failure must NOT look like a missing user", + resolver: &identity.ErrResolutionFailed{Identifier: handle, Reason: "PLC directory timeout"}, + assert: func(t *testing.T, err error) { + assert.NotErrorIs(t, err, ErrUserNotFound, + "reporting an outage as 404 hides it from 5xx alerting and misleads the caller") + var invalidHandle *InvalidHandleError + assert.NotErrorAs(t, err, &invalidHandle) + }, + }, + { + name: "unclassified resolver error stays unclassified", + resolver: errors.New("connection reset by peer"), + assert: func(t *testing.T, err error) { + assert.NotErrorIs(t, err, ErrUserNotFound) + assert.ErrorContains(t, err, "connection reset by peer") + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + mockRepo := new(MockUserRepository) + mockResolver := new(MockIdentityResolver) + + // Not indexed locally, so resolution falls through to the resolver. + mockRepo.On("GetByHandle", mock.Anything, mock.Anything).Return(nil, ErrUserNotFound) + mockResolver.On("ResolveHandle", mock.Anything, mock.Anything). + Return("", "", tt.resolver) + + service := NewUserService(mockRepo, mockResolver, "https://default.pds", nil, "") + + did, err := service.ResolveHandleToDID(context.Background(), handle) + assert.Empty(t, did) + assert.Error(t, err) + tt.assert(t, err) + }) + } +} + +// A handle already indexed locally must not touch the resolver at all. +func TestResolveHandleToDID_LocalHit(t *testing.T) { + mockRepo := new(MockUserRepository) + mockResolver := new(MockIdentityResolver) + + mockRepo.On("GetByHandle", mock.Anything, "known.example.com"). + Return(&User{DID: "did:plc:known", Handle: "known.example.com"}, nil) + + service := NewUserService(mockRepo, mockResolver, "https://default.pds", nil, "") + + did, err := service.ResolveHandleToDID(context.Background(), "known.example.com") + assert.NoError(t, err) + assert.Equal(t, "did:plc:known", did) + mockResolver.AssertNotCalled(t, "ResolveHandle", mock.Anything, mock.Anything) +} + +// A transient database error must not stop external resolution — the local +// lookup is only a cache. +func TestResolveHandleToDID_DatabaseErrorFallsThrough(t *testing.T) { + mockRepo := new(MockUserRepository) + mockResolver := new(MockIdentityResolver) + + mockRepo.On("GetByHandle", mock.Anything, mock.Anything). + Return(nil, fmt.Errorf("pq: too many connections")) + mockResolver.On("ResolveHandle", mock.Anything, mock.Anything). + Return("did:plc:resolved", "https://pds.example", nil) + + service := NewUserService(mockRepo, mockResolver, "https://default.pds", nil, "") + + did, err := service.ResolveHandleToDID(context.Background(), "someone.example.com") + assert.NoError(t, err) + assert.Equal(t, "did:plc:resolved", did) +} diff --git a/internal/core/users/service.go b/internal/core/users/service.go index c510d2f..f1d8542 100644 --- a/internal/core/users/service.go +++ b/internal/core/users/service.go @@ -224,6 +224,19 @@ func (s *userService) ResolveHandleToDID(ctx context.Context, handle string) (st // Slow path: use identity resolver for external DNS/HTTPS resolution did, _, err := s.identityResolver.ResolveHandle(ctx, handle) if err != nil { + // Translate the resolver's typed errors into this package's vocabulary so + // callers can tell "no such handle" from "resolution broke". Without this + // the two are indistinguishable at the API boundary and a nonexistent + // handle — the single most common failure of a public profile lookup — + // reads as a server fault. + var notFound *identity.ErrNotFound + var invalidIdentifier *identity.ErrInvalidIdentifier + switch { + case errors.As(err, ¬Found): + return "", fmt.Errorf("%w: %w", ErrUserNotFound, err) + case errors.As(err, &invalidIdentifier): + return "", &InvalidHandleError{Handle: handle, Reason: invalidIdentifier.Reason} + } return "", fmt.Errorf("failed to resolve handle %s: %w", handle, err) } diff --git a/internal/core/votes/cache.go b/internal/core/votes/cache.go index 690b341..1885fc3 100644 --- a/internal/core/votes/cache.go +++ b/internal/core/votes/cache.go @@ -170,7 +170,7 @@ func (c *VoteCache) fetchAllVotesFromPDS(ctx context.Context, pdsClient pds.Clie result, err := pdsClient.ListRecords(ctx, collection, pageSize, cursor) if err != nil { if pds.IsAuthError(err) { - return nil, ErrNotAuthorized + return nil, fmt.Errorf("%w: %w", ErrNotAuthorized, err) } return nil, fmt.Errorf("listRecords failed: %w", err) } diff --git a/internal/core/votes/errors.go b/internal/core/votes/errors.go index 89580d1..9926ea7 100644 --- a/internal/core/votes/errors.go +++ b/internal/core/votes/errors.go @@ -15,7 +15,17 @@ var ( // ErrVoteAlreadyExists indicates a vote already exists on this subject ErrVoteAlreadyExists = errors.New("vote already exists") - // ErrNotAuthorized indicates the user is not authorized to perform this action + // ErrNotAuthorized indicates the user is not authorized to perform this action. + // + // When the cause is a PDS auth failure, return it wrapped rather than bare: + // + // fmt.Errorf("%w: %w", ErrNotAuthorized, err) + // + // Returning the bare sentinel discards which auth failure it was, and the + // two need opposite responses: a PDS 401 means the session is dead and the + // client must sign in again, while a 403 means the session simply lacks the + // scope. Collapsed into one sentinel they both answer 403, so a client whose + // session expired retries forever with no way out but a manual sign-out. ErrNotAuthorized = errors.New("not authorized") // ErrBanned indicates the user is banned from the community diff --git a/internal/core/votes/service_impl.go b/internal/core/votes/service_impl.go index f6efa6b..dda538a 100644 --- a/internal/core/votes/service_impl.go +++ b/internal/core/votes/service_impl.go @@ -144,7 +144,7 @@ func (s *voteService) CreateVote(ctx context.Context, session *oauth.ClientSessi "voter", session.AccountDID, "rkey", existing.RKey) if pds.IsAuthError(err) { - return nil, ErrNotAuthorized + return nil, fmt.Errorf("%w: %w", ErrNotAuthorized, err) } return nil, fmt.Errorf("failed to delete vote: %w", err) } @@ -173,7 +173,7 @@ func (s *voteService) CreateVote(ctx context.Context, session *oauth.ClientSessi "voter", session.AccountDID, "rkey", existing.RKey) if pds.IsAuthError(err) { - return nil, ErrNotAuthorized + return nil, fmt.Errorf("%w: %w", ErrNotAuthorized, err) } return nil, fmt.Errorf("failed to delete existing vote: %w", err) } @@ -194,7 +194,7 @@ func (s *voteService) CreateVote(ctx context.Context, session *oauth.ClientSessi "subject", req.Subject.URI, "direction", req.Direction) if pds.IsAuthError(err) { - return nil, ErrNotAuthorized + return nil, fmt.Errorf("%w: %w", ErrNotAuthorized, err) } return nil, fmt.Errorf("failed to create vote: %w", err) } @@ -261,7 +261,7 @@ func (s *voteService) DeleteVote(ctx context.Context, session *oauth.ClientSessi "voter", session.AccountDID, "rkey", existing.RKey) if pds.IsAuthError(err) { - return ErrNotAuthorized + return fmt.Errorf("%w: %w", ErrNotAuthorized, err) } return fmt.Errorf("failed to delete vote: %w", err) } @@ -372,7 +372,7 @@ func (s *voteService) findExistingVoteFromPDS(ctx context.Context, pdsClient pds if err != nil { // Check for auth errors using typed errors if pds.IsAuthError(err) { - return nil, ErrNotAuthorized + return nil, fmt.Errorf("%w: %w", ErrNotAuthorized, err) } return nil, fmt.Errorf("listRecords failed: %w", err) } diff --git a/internal/db/postgres/aggregator_repo.go b/internal/db/postgres/aggregator_repo.go index 28908f2..9915df3 100644 --- a/internal/db/postgres/aggregator_repo.go +++ b/internal/db/postgres/aggregator_repo.go @@ -4,6 +4,7 @@ import ( "Coves/internal/core/aggregators" "context" "database/sql" + "errors" "fmt" "strings" "time" @@ -99,7 +100,7 @@ func (r *postgresAggregatorRepo) GetAggregator(ctx context.Context, did string) &recordCID, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, aggregators.ErrAggregatorNotFound } if err != nil { @@ -434,7 +435,7 @@ func (r *postgresAggregatorRepo) GetAuthorization(ctx context.Context, aggregato &recordCID, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, aggregators.ErrAuthorizationNotFound } if err != nil { @@ -487,7 +488,7 @@ func (r *postgresAggregatorRepo) GetAuthorizationByURI(ctx context.Context, reco &recordCID, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, aggregators.ErrAuthorizationNotFound } if err != nil { @@ -795,7 +796,7 @@ func (r *postgresAggregatorRepo) GetByAPIKeyHash(ctx context.Context, keyHash st &apiKeyRevokedAt, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, aggregators.ErrAggregatorNotFound } if err != nil { @@ -1025,7 +1026,7 @@ func (r *postgresAggregatorRepo) GetAggregatorCredentials(ctx context.Context, d &oauthDPoPPDSNonce, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, aggregators.ErrAggregatorNotFound } if err != nil { @@ -1119,7 +1120,7 @@ func (r *postgresAggregatorRepo) GetCredentialsByAPIKeyHash(ctx context.Context, &oauthDPoPPDSNonce, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, aggregators.ErrAPIKeyInvalid } if err != nil { diff --git a/internal/db/postgres/comment_repo.go b/internal/db/postgres/comment_repo.go index 73e2584..95c211a 100644 --- a/internal/db/postgres/comment_repo.go +++ b/internal/db/postgres/comment_repo.go @@ -5,6 +5,7 @@ import ( "context" "database/sql" "encoding/base64" + "errors" "fmt" "log" "strings" @@ -50,7 +51,7 @@ func (r *postgresCommentRepo) Create(ctx context.Context, comment *comments.Comm ).Scan(&comment.ID, &comment.IndexedAt) // ON CONFLICT DO NOTHING returns no rows if duplicate - this is OK (idempotent) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil // Comment already exists, no error for idempotency } @@ -130,7 +131,7 @@ func (r *postgresCommentRepo) Update(ctx context.Context, comment *comments.Comm &comment.ReplyCount, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return comments.ErrCommentNotFound } if err != nil { @@ -172,7 +173,7 @@ func (r *postgresCommentRepo) GetByURI(ctx context.Context, uri string) (*commen &comment.UpvoteCount, &comment.DownvoteCount, &comment.Score, &comment.ReplyCount, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, comments.ErrCommentNotFound } if err != nil { diff --git a/internal/db/postgres/community_repo.go b/internal/db/postgres/community_repo.go index 4ec872e..1cb955c 100644 --- a/internal/db/postgres/community_repo.go +++ b/internal/db/postgres/community_repo.go @@ -4,6 +4,7 @@ import ( "Coves/internal/core/communities" "context" "database/sql" + "errors" "fmt" "log" "strings" @@ -161,7 +162,7 @@ func (r *postgresCommunityRepo) GetByDID(ctx context.Context, did string) (*comm &recordURI, &recordCID, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, communities.ErrCommunityNotFound } if err != nil { @@ -225,7 +226,7 @@ func (r *postgresCommunityRepo) GetByHandle(ctx context.Context, handle string) &recordURI, &recordCID, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, communities.ErrCommunityNotFound } if err != nil { @@ -291,7 +292,7 @@ func (r *postgresCommunityRepo) Update(ctx context.Context, community *communiti community.PDSURL, ).Scan(&community.UpdatedAt) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, communities.ErrCommunityNotFound } if err != nil { @@ -317,7 +318,7 @@ func (r *postgresCommunityRepo) UpdateCredentials(ctx context.Context, did, acce var returnedDID string err := r.db.QueryRowContext(ctx, query, did, accessToken, refreshToken).Scan(&returnedDID) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return communities.ErrCommunityNotFound } if err != nil { diff --git a/internal/db/postgres/community_repo_blocks.go b/internal/db/postgres/community_repo_blocks.go index f1b172d..c962f4a 100644 --- a/internal/db/postgres/community_repo_blocks.go +++ b/internal/db/postgres/community_repo_blocks.go @@ -4,6 +4,7 @@ import ( "Coves/internal/core/communities" "context" "database/sql" + "errors" "fmt" "log" ) @@ -72,7 +73,7 @@ func (r *postgresCommunityRepo) GetBlock(ctx context.Context, userDID, community &block.RecordCID, ) if err != nil { - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, communities.ErrBlockNotFound } return nil, fmt.Errorf("failed to get block: %w", err) @@ -99,7 +100,7 @@ func (r *postgresCommunityRepo) GetBlockByURI(ctx context.Context, recordURI str &block.RecordCID, ) if err != nil { - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, communities.ErrBlockNotFound } return nil, fmt.Errorf("failed to get block by URI: %w", err) diff --git a/internal/db/postgres/community_repo_memberships.go b/internal/db/postgres/community_repo_memberships.go index 562c9db..7a2f21f 100644 --- a/internal/db/postgres/community_repo_memberships.go +++ b/internal/db/postgres/community_repo_memberships.go @@ -4,6 +4,7 @@ import ( "Coves/internal/core/communities" "context" "database/sql" + "errors" "fmt" "log" "strings" @@ -63,7 +64,7 @@ func (r *postgresCommunityRepo) GetMembership(ctx context.Context, userDID, comm &membership.IsModerator, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, communities.ErrMembershipNotFound } if err != nil { @@ -95,7 +96,7 @@ func (r *postgresCommunityRepo) UpdateMembership(ctx context.Context, membership membership.IsModerator, ).Scan(&membership.LastActiveAt) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, communities.ErrMembershipNotFound } if err != nil { diff --git a/internal/db/postgres/community_repo_subscriptions.go b/internal/db/postgres/community_repo_subscriptions.go index fb03697..d6fb0d3 100644 --- a/internal/db/postgres/community_repo_subscriptions.go +++ b/internal/db/postgres/community_repo_subscriptions.go @@ -4,6 +4,7 @@ import ( "Coves/internal/core/communities" "context" "database/sql" + "errors" "fmt" "log" "strings" @@ -203,7 +204,7 @@ func (r *postgresCommunityRepo) GetSubscription(ctx context.Context, userDID, co &subscription.ContentVisibility, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, communities.ErrSubscriptionNotFound } if err != nil { @@ -237,7 +238,7 @@ func (r *postgresCommunityRepo) GetSubscriptionByURI(ctx context.Context, record &subscription.ContentVisibility, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, communities.ErrSubscriptionNotFound } if err != nil { diff --git a/internal/db/postgres/community_suggestion_repo.go b/internal/db/postgres/community_suggestion_repo.go index 47a702b..fd5981a 100644 --- a/internal/db/postgres/community_suggestion_repo.go +++ b/internal/db/postgres/community_suggestion_repo.go @@ -4,6 +4,7 @@ import ( "Coves/internal/core/communitysuggestions" "context" "database/sql" + "errors" "fmt" "log/slog" "strings" @@ -88,7 +89,7 @@ func (r *postgresCommunitySuggestionRepo) GetByID(ctx context.Context, id int64) &suggestion.VoteCount, &suggestion.CreatedAt, &suggestion.UpdatedAt, ) if err != nil { - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, communitysuggestions.ErrSuggestionNotFound } return nil, fmt.Errorf("failed to get community suggestion by ID: %w", err) @@ -319,7 +320,7 @@ func (r *postgresCommunitySuggestionRepo) DeleteVote(ctx context.Context, sugges var deletedValue int err = tx.QueryRowContext(ctx, deleteQuery, suggestionID, voterDID).Scan(&deletedValue) if err != nil { - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return 0, communitysuggestions.ErrVoteNotFound } return 0, fmt.Errorf("failed to delete vote: %w", err) @@ -480,7 +481,7 @@ func (r *postgresCommunitySuggestionRepo) GetVote(ctx context.Context, suggestio &vote.Value, &vote.CreatedAt, ) if err != nil { - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, communitysuggestions.ErrVoteNotFound } return nil, fmt.Errorf("failed to get vote: %w", err) diff --git a/internal/db/postgres/post_repo.go b/internal/db/postgres/post_repo.go index ef411e1..eb27f25 100644 --- a/internal/db/postgres/post_repo.go +++ b/internal/db/postgres/post_repo.go @@ -5,6 +5,7 @@ import ( "database/sql" "encoding/base64" "encoding/json" + "errors" "fmt" "log/slog" "strings" @@ -134,7 +135,7 @@ func (r *postgresPostRepo) GetByURI(ctx context.Context, uri string) (*posts.Pos &post.UpvoteCount, &post.DownvoteCount, &post.Score, &post.CommentCount, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, posts.ErrNotFound } if err != nil { diff --git a/internal/db/postgres/user_repo.go b/internal/db/postgres/user_repo.go index 8a0f13d..107ff70 100644 --- a/internal/db/postgres/user_repo.go +++ b/internal/db/postgres/user_repo.go @@ -4,6 +4,7 @@ import ( "Coves/internal/core/users" "context" "database/sql" + "errors" "fmt" "log/slog" "strings" @@ -55,7 +56,7 @@ func (r *postgresUserRepo) GetByDID(ctx context.Context, did string) (*users.Use Scan(&user.DID, &user.Handle, &user.PDSURL, &user.CreatedAt, &user.UpdatedAt, &displayName, &bio, &avatarCID, &bannerCID) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, users.ErrUserNotFound } if err != nil { @@ -80,7 +81,7 @@ func (r *postgresUserRepo) GetByHandle(ctx context.Context, handle string) (*use Scan(&user.DID, &user.Handle, &user.PDSURL, &user.CreatedAt, &user.UpdatedAt, &displayName, &bio, &avatarCID, &bannerCID) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, users.ErrUserNotFound } if err != nil { @@ -109,7 +110,7 @@ func (r *postgresUserRepo) UpdateHandle(ctx context.Context, did, newHandle stri Scan(&user.DID, &user.Handle, &user.PDSURL, &user.CreatedAt, &user.UpdatedAt, &displayName, &bio, &avatarCID, &bannerCID) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, users.ErrUserNotFound } if err != nil { @@ -370,7 +371,7 @@ func (r *postgresUserRepo) UpdateProfile(ctx context.Context, did string, input Scan(&user.DID, &user.Handle, &user.PDSURL, &user.CreatedAt, &user.UpdatedAt, &displayNameVal, &bioVal, &avatarCIDVal, &bannerCIDVal) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, users.ErrUserNotFound } if err != nil { diff --git a/internal/db/postgres/vote_repo.go b/internal/db/postgres/vote_repo.go index 278c49d..0c2b3aa 100644 --- a/internal/db/postgres/vote_repo.go +++ b/internal/db/postgres/vote_repo.go @@ -4,6 +4,7 @@ import ( "Coves/internal/core/votes" "context" "database/sql" + "errors" "fmt" "strings" ) @@ -43,7 +44,7 @@ func (r *postgresVoteRepo) Create(ctx context.Context, vote *votes.Vote) error { ).Scan(&vote.ID, &vote.IndexedAt) // ON CONFLICT DO NOTHING returns no rows if duplicate - this is OK (idempotent) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil // Vote already exists, no error for idempotency } @@ -85,7 +86,7 @@ func (r *postgresVoteRepo) GetByURI(ctx context.Context, uri string) (*votes.Vot &vote.CreatedAt, &vote.IndexedAt, &vote.DeletedAt, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, votes.ErrVoteNotFound } if err != nil { @@ -115,7 +116,7 @@ func (r *postgresVoteRepo) GetByVoterAndSubject(ctx context.Context, voterDID, s &vote.CreatedAt, &vote.IndexedAt, &vote.DeletedAt, ) - if err == sql.ErrNoRows { + if errors.Is(err, sql.ErrNoRows) { return nil, votes.ErrVoteNotFound } if err != nil {