fix(auth): enforce private instance access boundaries
- Resolve Gateway procedures from matched HTTP bindings before authorization.\n- Disable anonymous RSS on private instances.\n- Limit share-token access to the shared memo and its attachments.
This commit is contained in:
parent
0d2cbd4f5a
commit
415a3ec73d
8 changed files with 637 additions and 28 deletions
189
server/router/api/v1/gateway_route_resolver.go
Normal file
189
server/router/api/v1/gateway_route_resolver.go
Normal file
|
|
@ -0,0 +1,189 @@
|
|||
package v1
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/grpc-ecosystem/grpc-gateway/v2/runtime"
|
||||
"github.com/pkg/errors"
|
||||
"google.golang.org/genproto/googleapis/api/annotations"
|
||||
"google.golang.org/protobuf/proto"
|
||||
"google.golang.org/protobuf/reflect/protoreflect"
|
||||
"google.golang.org/protobuf/reflect/protoregistry"
|
||||
)
|
||||
|
||||
// apiPackage is the proto package whose HTTP bindings the gateway serves.
|
||||
const apiPackage = "memos.api.v1"
|
||||
|
||||
type gatewayRouteKey struct {
|
||||
httpMethod string
|
||||
pattern string
|
||||
}
|
||||
|
||||
// gatewayRouteResolver maps the HTTP pattern already matched by grpc-gateway
|
||||
// back to the gRPC procedure that owns the binding.
|
||||
//
|
||||
// runtime.RPCMethod is not available to middleware because the generated
|
||||
// handler sets it after middleware runs. The ServeMux does, however, put the
|
||||
// exact matched runtime.Pattern in the request context before invoking the
|
||||
// middleware. Using that pattern avoids independently reimplementing the
|
||||
// gateway's path matching, escaping, verb, and route-ordering semantics.
|
||||
type gatewayRouteResolver struct {
|
||||
procedures map[gatewayRouteKey]string
|
||||
}
|
||||
|
||||
// newGatewayRouteResolver builds a lookup table from the registered protos.
|
||||
func newGatewayRouteResolver() (*gatewayRouteResolver, error) {
|
||||
resolver := &gatewayRouteResolver{procedures: map[gatewayRouteKey]string{}}
|
||||
|
||||
var rangeErr error
|
||||
protoregistry.GlobalFiles.RangeFiles(func(fd protoreflect.FileDescriptor) bool {
|
||||
if string(fd.Package()) != apiPackage {
|
||||
return true
|
||||
}
|
||||
services := fd.Services()
|
||||
for i := range services.Len() {
|
||||
service := services.Get(i)
|
||||
methods := service.Methods()
|
||||
for j := range methods.Len() {
|
||||
method := methods.Get(j)
|
||||
rule, ok := proto.GetExtension(method.Options(), annotations.E_Http).(*annotations.HttpRule)
|
||||
if !ok || rule == nil {
|
||||
continue
|
||||
}
|
||||
procedure := "/" + string(service.FullName()) + "/" + string(method.Name())
|
||||
if err := resolver.addRule(procedure, rule); err != nil {
|
||||
rangeErr = err
|
||||
return false
|
||||
}
|
||||
}
|
||||
}
|
||||
return true
|
||||
})
|
||||
if rangeErr != nil {
|
||||
return nil, rangeErr
|
||||
}
|
||||
if len(resolver.procedures) == 0 {
|
||||
return nil, errors.Errorf("no HTTP bindings found for proto package %s", apiPackage)
|
||||
}
|
||||
return resolver, nil
|
||||
}
|
||||
|
||||
// addRule records a rule and each of its additional bindings.
|
||||
func (r *gatewayRouteResolver) addRule(procedure string, rule *annotations.HttpRule) error {
|
||||
if err := r.addBinding(procedure, rule); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, binding := range rule.GetAdditionalBindings() {
|
||||
if err := r.addBinding(procedure, binding); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *gatewayRouteResolver) addBinding(procedure string, rule *annotations.HttpRule) error {
|
||||
httpMethod, template := httpRuleMethodAndTemplate(rule)
|
||||
if httpMethod == "" || template == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
pattern, err := canonicalGatewayPattern(template)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to canonicalize HTTP template %q for %s", template, procedure)
|
||||
}
|
||||
key := gatewayRouteKey{httpMethod: httpMethod, pattern: pattern}
|
||||
if existing, ok := r.procedures[key]; ok && existing != procedure {
|
||||
return errors.Errorf("%s %s is bound to both %s and %s", httpMethod, pattern, existing, procedure)
|
||||
}
|
||||
r.procedures[key] = procedure
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *gatewayRouteResolver) resolve(httpMethod, pattern string) (string, bool) {
|
||||
procedure, ok := r.procedures[gatewayRouteKey{httpMethod: httpMethod, pattern: pattern}]
|
||||
return procedure, ok
|
||||
}
|
||||
|
||||
// resolveRequest returns the procedure for the handler grpc-gateway selected.
|
||||
func (r *gatewayRouteResolver) resolveRequest(request *http.Request) (string, bool) {
|
||||
pattern, ok := runtime.HTTPPattern(request.Context())
|
||||
if !ok {
|
||||
return "", false
|
||||
}
|
||||
|
||||
matchedPattern := pattern.String()
|
||||
if procedure, ok := r.resolve(request.Method, matchedPattern); ok {
|
||||
return procedure, true
|
||||
}
|
||||
|
||||
// grpc-gateway supports submitting a GET binding as a form POST when no
|
||||
// POST binding matched the path. The selected pattern is still the GET
|
||||
// handler's pattern, while request.Method remains POST.
|
||||
if request.Method == http.MethodPost && request.Header.Get("Content-Type") == "application/x-www-form-urlencoded" {
|
||||
return r.resolve(http.MethodGet, matchedPattern)
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
// httpRuleMethodAndTemplate extracts the HTTP method and path template from a rule.
|
||||
func httpRuleMethodAndTemplate(rule *annotations.HttpRule) (string, string) {
|
||||
switch pattern := rule.GetPattern().(type) {
|
||||
case *annotations.HttpRule_Get:
|
||||
return http.MethodGet, pattern.Get
|
||||
case *annotations.HttpRule_Post:
|
||||
return http.MethodPost, pattern.Post
|
||||
case *annotations.HttpRule_Put:
|
||||
return http.MethodPut, pattern.Put
|
||||
case *annotations.HttpRule_Patch:
|
||||
return http.MethodPatch, pattern.Patch
|
||||
case *annotations.HttpRule_Delete:
|
||||
return http.MethodDelete, pattern.Delete
|
||||
case *annotations.HttpRule_Custom:
|
||||
return strings.ToUpper(pattern.Custom.GetKind()), pattern.Custom.GetPath()
|
||||
default:
|
||||
return "", ""
|
||||
}
|
||||
}
|
||||
|
||||
// canonicalGatewayPattern converts a google.api.http template to the normalized
|
||||
// form returned by runtime.Pattern.String. In particular, grpc-gateway expands
|
||||
// a bare "{field}" variable to "{field=*}".
|
||||
func canonicalGatewayPattern(template string) (string, error) {
|
||||
if !strings.HasPrefix(template, "/") {
|
||||
return "", errors.New("template must start with '/'")
|
||||
}
|
||||
|
||||
var builder strings.Builder
|
||||
for index := 0; index < len(template); {
|
||||
switch template[index] {
|
||||
case '{':
|
||||
relativeEnd := strings.IndexByte(template[index+1:], '}')
|
||||
if relativeEnd < 0 {
|
||||
return "", errors.New("unclosed variable")
|
||||
}
|
||||
end := index + 1 + relativeEnd
|
||||
variable := template[index+1 : end]
|
||||
if variable == "" || strings.ContainsAny(variable, "{}") {
|
||||
return "", errors.Errorf("invalid variable %q", variable)
|
||||
}
|
||||
field, subtemplate, hasSubtemplate := strings.Cut(variable, "=")
|
||||
if field == "" || (hasSubtemplate && subtemplate == "") {
|
||||
return "", errors.Errorf("invalid variable %q", variable)
|
||||
}
|
||||
if !hasSubtemplate {
|
||||
variable += "=*"
|
||||
}
|
||||
builder.WriteByte('{')
|
||||
builder.WriteString(variable)
|
||||
builder.WriteByte('}')
|
||||
index = end + 1
|
||||
case '}':
|
||||
return "", errors.New("unexpected closing brace")
|
||||
default:
|
||||
builder.WriteByte(template[index])
|
||||
index++
|
||||
}
|
||||
}
|
||||
return builder.String(), nil
|
||||
}
|
||||
318
server/router/api/v1/gateway_route_resolver_test.go
Normal file
318
server/router/api/v1/gateway_route_resolver_test.go
Normal file
|
|
@ -0,0 +1,318 @@
|
|||
package v1
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/grpc-ecosystem/grpc-gateway/v2/runtime"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/genproto/googleapis/api/annotations"
|
||||
"google.golang.org/protobuf/proto"
|
||||
"google.golang.org/protobuf/reflect/protoreflect"
|
||||
"google.golang.org/protobuf/reflect/protoregistry"
|
||||
)
|
||||
|
||||
func TestCanonicalGatewayPattern(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
template string
|
||||
want string
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "literal",
|
||||
template: "/api/v1/users",
|
||||
want: "/api/v1/users",
|
||||
},
|
||||
{
|
||||
name: "bare variable",
|
||||
template: "/api/v1/shares/{share_token}/memo",
|
||||
want: "/api/v1/shares/{share_token=*}/memo",
|
||||
},
|
||||
{
|
||||
name: "resource name variable",
|
||||
template: "/api/v1/{name=users/*}",
|
||||
want: "/api/v1/{name=users/*}",
|
||||
},
|
||||
{
|
||||
name: "variable with verb",
|
||||
template: "/api/v1/{name=users/*}:getStats",
|
||||
want: "/api/v1/{name=users/*}:getStats",
|
||||
},
|
||||
{
|
||||
name: "missing leading slash",
|
||||
template: "api/v1/users",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "unclosed variable",
|
||||
template: "/api/v1/{name",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "unexpected closing brace",
|
||||
template: "/api/v1/name}",
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
got, err := canonicalGatewayPattern(test.template)
|
||||
if test.wantErr {
|
||||
require.Error(t, err)
|
||||
return
|
||||
}
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, test.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestGatewayRouteResolverResolvesKnownRoutes pins the procedures that carry
|
||||
// access-control decisions. The request is dispatched through a real ServeMux
|
||||
// so the resolver consumes the exact runtime.Pattern seen by middleware.
|
||||
func TestGatewayRouteResolverResolvesKnownRoutes(t *testing.T) {
|
||||
resolver, err := newGatewayRouteResolver()
|
||||
require.NoError(t, err)
|
||||
|
||||
tests := []struct {
|
||||
httpMethod string
|
||||
template string
|
||||
path string
|
||||
procedure string
|
||||
}{
|
||||
{http.MethodPost, "/api/v1/auth/signin", "/api/v1/auth/signin", "/memos.api.v1.AuthService/SignIn"},
|
||||
{http.MethodGet, "/api/v1/instance/profile", "/api/v1/instance/profile", "/memos.api.v1.InstanceService/GetInstanceProfile"},
|
||||
{http.MethodPost, "/api/v1/users", "/api/v1/users", "/memos.api.v1.UserService/CreateUser"},
|
||||
{http.MethodGet, "/api/v1/users", "/api/v1/users", "/memos.api.v1.UserService/ListUsers"},
|
||||
{http.MethodGet, "/api/v1/{name=users/*}", "/api/v1/users/1", "/memos.api.v1.UserService/GetUser"},
|
||||
{http.MethodGet, "/api/v1/memos", "/api/v1/memos", "/memos.api.v1.MemoService/ListMemos"},
|
||||
{http.MethodGet, "/api/v1/{name=memos/*}", "/api/v1/memos/abc", "/memos.api.v1.MemoService/GetMemo"},
|
||||
{http.MethodPost, "/api/v1/memos", "/api/v1/memos", "/memos.api.v1.MemoService/CreateMemo"},
|
||||
{
|
||||
http.MethodGet,
|
||||
"/api/v1/shares/{share_token}/memo",
|
||||
"/api/v1/shares/token:with-colon/memo",
|
||||
"/memos.api.v1.MemoService/GetSharedMemo",
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.httpMethod+" "+test.path, func(t *testing.T) {
|
||||
procedure, ok := resolveThroughGateway(t, resolver, test.httpMethod, test.template, test.httpMethod, test.path, "")
|
||||
require.True(t, ok)
|
||||
require.Equal(t, test.procedure, procedure)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGatewayRouteResolverResolvesFormPostFallback(t *testing.T) {
|
||||
resolver, err := newGatewayRouteResolver()
|
||||
require.NoError(t, err)
|
||||
|
||||
t.Run("falls back to public GET binding", func(t *testing.T) {
|
||||
procedure, ok := resolveThroughGateway(
|
||||
t,
|
||||
resolver,
|
||||
http.MethodGet,
|
||||
"/api/v1/instance/profile",
|
||||
http.MethodPost,
|
||||
"/api/v1/instance/profile",
|
||||
"application/x-www-form-urlencoded",
|
||||
)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, "/memos.api.v1.InstanceService/GetInstanceProfile", procedure)
|
||||
})
|
||||
|
||||
t.Run("uses matching POST binding", func(t *testing.T) {
|
||||
procedure, ok := resolveThroughGateway(
|
||||
t,
|
||||
resolver,
|
||||
http.MethodPost,
|
||||
"/api/v1/memos",
|
||||
http.MethodPost,
|
||||
"/api/v1/memos",
|
||||
"application/x-www-form-urlencoded",
|
||||
)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, "/memos.api.v1.MemoService/CreateMemo", procedure)
|
||||
})
|
||||
}
|
||||
|
||||
func TestGatewayRouteResolverRequiresMatchedPattern(t *testing.T) {
|
||||
resolver, err := newGatewayRouteResolver()
|
||||
require.NoError(t, err)
|
||||
|
||||
request := httptest.NewRequest(http.MethodGet, "/api/v1/memos", nil)
|
||||
_, ok := resolver.resolveRequest(request)
|
||||
require.False(t, ok)
|
||||
}
|
||||
|
||||
// TestGatewayRouteResolverCoversEveryBinding walks every google.api.http binding
|
||||
// in the API package, dispatches a concrete request through grpc-gateway, and
|
||||
// requires the matched pattern to map back to the same procedure.
|
||||
func TestGatewayRouteResolverCoversEveryBinding(t *testing.T) {
|
||||
resolver, err := newGatewayRouteResolver()
|
||||
require.NoError(t, err)
|
||||
|
||||
checked := 0
|
||||
protoregistry.GlobalFiles.RangeFiles(func(fd protoreflect.FileDescriptor) bool {
|
||||
if string(fd.Package()) != apiPackage {
|
||||
return true
|
||||
}
|
||||
services := fd.Services()
|
||||
for i := range services.Len() {
|
||||
service := services.Get(i)
|
||||
methods := service.Methods()
|
||||
for j := range methods.Len() {
|
||||
method := methods.Get(j)
|
||||
rule, ok := proto.GetExtension(method.Options(), annotations.E_Http).(*annotations.HttpRule)
|
||||
if !ok || rule == nil {
|
||||
continue
|
||||
}
|
||||
procedure := "/" + string(service.FullName()) + "/" + string(method.Name())
|
||||
|
||||
for _, binding := range append([]*annotations.HttpRule{rule}, rule.GetAdditionalBindings()...) {
|
||||
httpMethod, template := httpRuleMethodAndTemplate(binding)
|
||||
if httpMethod == "" || template == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
path := synthesizeGatewayPath(template)
|
||||
resolved, found := resolveThroughGateway(t, resolver, httpMethod, template, httpMethod, path, "")
|
||||
require.True(t, found, "%s %s (from %q) should resolve", httpMethod, path, template)
|
||||
require.Equal(t, procedure, resolved,
|
||||
"%s %s (from template %q) resolved to the wrong procedure", httpMethod, path, template)
|
||||
checked++
|
||||
}
|
||||
}
|
||||
}
|
||||
return true
|
||||
})
|
||||
|
||||
require.Positive(t, checked, "should have checked at least one binding")
|
||||
t.Logf("verified %d HTTP bindings", checked)
|
||||
}
|
||||
|
||||
func resolveThroughGateway(
|
||||
t *testing.T,
|
||||
resolver *gatewayRouteResolver,
|
||||
bindingMethod string,
|
||||
template string,
|
||||
requestMethod string,
|
||||
path string,
|
||||
contentType string,
|
||||
) (string, bool) {
|
||||
t.Helper()
|
||||
|
||||
var (
|
||||
called bool
|
||||
procedure string
|
||||
resolved bool
|
||||
)
|
||||
mux := runtime.NewServeMux()
|
||||
require.NoError(t, mux.HandlePath(bindingMethod, template, func(_ http.ResponseWriter, request *http.Request, _ map[string]string) {
|
||||
called = true
|
||||
procedure, resolved = resolver.resolveRequest(request)
|
||||
}))
|
||||
|
||||
request := httptest.NewRequest(requestMethod, path, nil)
|
||||
if contentType != "" {
|
||||
request.Header.Set("Content-Type", contentType)
|
||||
}
|
||||
mux.ServeHTTP(httptest.NewRecorder(), request)
|
||||
require.True(t, called, "grpc-gateway should match %s %s against %q", requestMethod, path, template)
|
||||
return procedure, resolved
|
||||
}
|
||||
|
||||
// synthesizeGatewayPath turns a path template into a concrete path by replacing
|
||||
// variables and wildcards with placeholder segments.
|
||||
func synthesizeGatewayPath(template string) string {
|
||||
path, verb := splitGatewayTemplateVerb(template)
|
||||
|
||||
segments := splitGatewayPathSegments(path)
|
||||
built := make([]string, 0, len(segments))
|
||||
for _, segment := range segments {
|
||||
switch {
|
||||
case segment == "*":
|
||||
built = append(built, "one")
|
||||
case segment == "**":
|
||||
built = append(built, "one", "two")
|
||||
case strings.HasPrefix(segment, "{") && strings.HasSuffix(segment, "}"):
|
||||
built = append(built, synthesizeGatewayVariable(segment)...)
|
||||
default:
|
||||
built = append(built, segment)
|
||||
}
|
||||
}
|
||||
|
||||
result := "/" + strings.Join(built, "/")
|
||||
if verb != "" {
|
||||
result += ":" + verb
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func splitGatewayPathSegments(path string) []string {
|
||||
trimmed := strings.TrimPrefix(path, "/")
|
||||
if trimmed == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
var segments []string
|
||||
depth, start := 0, 0
|
||||
for index := 0; index < len(trimmed); index++ {
|
||||
switch trimmed[index] {
|
||||
case '{':
|
||||
depth++
|
||||
case '}':
|
||||
depth--
|
||||
case '/':
|
||||
if depth == 0 {
|
||||
segments = append(segments, trimmed[start:index])
|
||||
start = index + 1
|
||||
}
|
||||
}
|
||||
}
|
||||
return append(segments, trimmed[start:])
|
||||
}
|
||||
|
||||
func splitGatewayTemplateVerb(template string) (string, string) {
|
||||
depth := 0
|
||||
for index := len(template) - 1; index >= 0; index-- {
|
||||
switch template[index] {
|
||||
case '}':
|
||||
depth++
|
||||
case '{':
|
||||
depth--
|
||||
case ':':
|
||||
if depth == 0 {
|
||||
return template[:index], template[index+1:]
|
||||
}
|
||||
}
|
||||
}
|
||||
return template, ""
|
||||
}
|
||||
|
||||
func synthesizeGatewayVariable(segment string) []string {
|
||||
inner := strings.TrimSuffix(strings.TrimPrefix(segment, "{"), "}")
|
||||
_, subtemplate, hasSubtemplate := strings.Cut(inner, "=")
|
||||
if !hasSubtemplate || subtemplate == "" || subtemplate == "*" {
|
||||
return []string{"one"}
|
||||
}
|
||||
|
||||
var built []string
|
||||
for _, part := range strings.Split(subtemplate, "/") {
|
||||
switch part {
|
||||
case "*":
|
||||
built = append(built, "one")
|
||||
case "**":
|
||||
built = append(built, "one", "two")
|
||||
default:
|
||||
built = append(built, part)
|
||||
}
|
||||
}
|
||||
return built
|
||||
}
|
||||
|
|
@ -176,18 +176,17 @@ func (s *APIV1Service) GetSharedMemo(ctx context.Context, request *v1pb.GetShare
|
|||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to list attachments")
|
||||
}
|
||||
relations, err := s.batchConvertMemoRelations(ctx, []*store.Memo{memo}, true)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to load memo relations")
|
||||
}
|
||||
|
||||
memoMessage, err := s.convertMemoFromStore(ctx, memo, reactions, attachments, relations[memo.ID])
|
||||
memoMessage, err := s.convertMemoFromStore(ctx, memo, reactions, attachments, nil)
|
||||
if err != nil {
|
||||
if stderrors.Is(err, errMemoCreatorNotFound) {
|
||||
return nil, status.Errorf(codes.NotFound, "not found")
|
||||
}
|
||||
return nil, errors.Wrap(err, "failed to convert memo")
|
||||
}
|
||||
// A share token grants access to this memo only, not to its surrounding
|
||||
// conversation or relation graph.
|
||||
memoMessage.Parent = nil
|
||||
return memoMessage, nil
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -110,6 +110,54 @@ func TestGetSharedMemo_IncludesReactions(t *testing.T) {
|
|||
require.Equal(t, memo.Name, sharedMemo.Reactions[0].ContentId)
|
||||
}
|
||||
|
||||
func TestGetSharedMemo_ExcludesParentAndRelations(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
ts := NewTestService(t)
|
||||
defer ts.Cleanup()
|
||||
|
||||
user, err := ts.CreateRegularUser(ctx, "share-single-memo")
|
||||
require.NoError(t, err)
|
||||
userCtx := ts.CreateUserContext(ctx, user.ID)
|
||||
|
||||
parent, err := ts.Service.CreateMemo(userCtx, &apiv1.CreateMemoRequest{
|
||||
Memo: &apiv1.Memo{
|
||||
Content: "parent must not be shared",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
comment, err := ts.Service.CreateMemoComment(userCtx, &apiv1.CreateMemoCommentRequest{
|
||||
Name: parent.Name,
|
||||
Comment: &apiv1.Memo{
|
||||
Content: "only this memo is shared",
|
||||
Visibility: apiv1.Visibility_PRIVATE,
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, comment.Relations)
|
||||
|
||||
regularMemo, err := ts.Service.GetMemo(userCtx, &apiv1.GetMemoRequest{Name: comment.Name})
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, regularMemo.GetParent())
|
||||
require.NotEmpty(t, regularMemo.Relations)
|
||||
|
||||
share, err := ts.Service.CreateMemoShare(userCtx, &apiv1.CreateMemoShareRequest{
|
||||
Parent: comment.Name,
|
||||
MemoShare: &apiv1.MemoShare{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
shareToken := share.Name[strings.LastIndex(share.Name, "/")+1:]
|
||||
sharedMemo, err := ts.Service.GetSharedMemo(ctx, &apiv1.GetSharedMemoRequest{
|
||||
ShareToken: shareToken,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, sharedMemo.GetParent())
|
||||
require.Empty(t, sharedMemo.Relations)
|
||||
}
|
||||
|
||||
func TestGetSharedMemo_SkipsReactionsWithMissingCreators(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ import (
|
|||
"connectrpc.com/connect"
|
||||
"github.com/grpc-ecosystem/grpc-gateway/v2/runtime"
|
||||
"github.com/labstack/echo/v5"
|
||||
"github.com/pkg/errors"
|
||||
"golang.org/x/sync/semaphore"
|
||||
|
||||
"github.com/usememos/memos/internal/markdown"
|
||||
|
|
@ -69,21 +70,31 @@ func (s *APIV1Service) RegisterGateway(ctx context.Context, echoServer *echo.Ech
|
|||
// Shared authorizer: one source of truth for authentication and anonymous-access
|
||||
// policy, used by both the gRPC-Gateway middleware and the Connect interceptor.
|
||||
authorizer := NewAuthorizer(s.Store, s.Secret, s.Profile)
|
||||
|
||||
// grpc-gateway does not hand the matched procedure to middleware:
|
||||
// runtime.RPCMethod is only populated by the generated handler, which runs
|
||||
// after middleware. Resolve the procedure from the proto HTTP bindings
|
||||
// instead so the policy check actually runs on this transport.
|
||||
routeResolver, err := newGatewayRouteResolver()
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "failed to build gateway route resolver")
|
||||
}
|
||||
|
||||
gatewayAuthMiddleware := func(next runtime.HandlerFunc) runtime.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request, pathParams map[string]string) {
|
||||
ctx := r.Context()
|
||||
|
||||
// The RPC method name is set by grpc-gateway after routing. When it can't be
|
||||
// determined, skip the policy check and let the service layer handle visibility.
|
||||
rpcMethod, ok := runtime.RPCMethod(ctx)
|
||||
authHeader := r.Header.Get("Authorization")
|
||||
|
||||
result := authorizer.Authenticate(ctx, authHeader)
|
||||
if ok {
|
||||
if err := authorizer.CheckAccess(ctx, rpcMethod, result); err != nil {
|
||||
http.Error(w, `{"code": 16, "message": "authentication required"}`, http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
// An unresolved path yields an empty procedure, which CheckAccess
|
||||
// treats as protected: authenticated callers pass and anonymous ones
|
||||
// are refused. Failing closed keeps a routing gap from becoming an
|
||||
// access-control gap.
|
||||
procedure, _ := routeResolver.resolveRequest(r)
|
||||
if err := authorizer.CheckAccess(ctx, procedure, result); err != nil {
|
||||
http.Error(w, `{"code": 16, "message": "authentication required"}`, http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
// Apply the identity to the context (no-op for permitted anonymous requests).
|
||||
|
|
|
|||
|
|
@ -70,6 +70,10 @@ func (s *RSSService) RegisterRoutes(g *echo.Group) {
|
|||
}
|
||||
|
||||
func (s *RSSService) GetExploreRSS(c *echo.Context) error {
|
||||
if s.Profile == nil || !s.Profile.AllowAnonymous() {
|
||||
return echo.NewHTTPError(http.StatusNotFound, "RSS is unavailable")
|
||||
}
|
||||
|
||||
ctx := c.Request().Context()
|
||||
cacheKey := "explore"
|
||||
|
||||
|
|
@ -109,6 +113,10 @@ func (s *RSSService) GetExploreRSS(c *echo.Context) error {
|
|||
}
|
||||
|
||||
func (s *RSSService) GetUserRSS(c *echo.Context) error {
|
||||
if s.Profile == nil || !s.Profile.AllowAnonymous() {
|
||||
return echo.NewHTTPError(http.StatusNotFound, "RSS is unavailable")
|
||||
}
|
||||
|
||||
ctx := c.Request().Context()
|
||||
username := c.Param("username")
|
||||
cacheKey := "user:" + username
|
||||
|
|
|
|||
|
|
@ -51,7 +51,7 @@ func TestPublicRSSExcludesComments(t *testing.T) {
|
|||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
service := NewRSSService(&profile.Profile{}, stores, markdown.NewService())
|
||||
service := NewRSSService(&profile.Profile{InstanceURL: "https://memos.example.com"}, stores, markdown.NewService())
|
||||
|
||||
exploreRSS := renderRSS(t, service, "/explore/rss.xml", "")
|
||||
require.Contains(t, exploreRSS, "public parent should stay in rss")
|
||||
|
|
@ -62,6 +62,40 @@ func TestPublicRSSExcludesComments(t *testing.T) {
|
|||
require.NotContains(t, userRSS, "public comment should not be in rss")
|
||||
}
|
||||
|
||||
func TestPrivateInstanceDisablesRSS(t *testing.T) {
|
||||
service := NewRSSService(&profile.Profile{}, nil, nil)
|
||||
|
||||
for _, test := range []struct {
|
||||
name string
|
||||
target string
|
||||
username string
|
||||
}{
|
||||
{name: "explore", target: "/explore/rss.xml"},
|
||||
{name: "user", target: "/u/alice/rss.xml", username: "alice"},
|
||||
} {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
e := echo.New()
|
||||
req := httptest.NewRequest(http.MethodGet, test.target, strings.NewReader(""))
|
||||
rec := httptest.NewRecorder()
|
||||
c := e.NewContext(req, rec)
|
||||
if test.username != "" {
|
||||
c.SetPathValues(echo.PathValues{{Name: "username", Value: test.username}})
|
||||
}
|
||||
|
||||
var err error
|
||||
if test.username == "" {
|
||||
err = service.GetExploreRSS(c)
|
||||
} else {
|
||||
err = service.GetUserRSS(c)
|
||||
}
|
||||
|
||||
var httpError *echo.HTTPError
|
||||
require.ErrorAs(t, err, &httpError)
|
||||
require.Equal(t, http.StatusNotFound, httpError.Code)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func renderRSS(t *testing.T, service *RSSService, target string, username string) string {
|
||||
t.Helper()
|
||||
|
||||
|
|
|
|||
|
|
@ -49,7 +49,7 @@ const MemoDetail = () => {
|
|||
});
|
||||
|
||||
const { data: parentMemo } = useMemo(memo?.parent || "", {
|
||||
enabled: !!memo?.parent,
|
||||
enabled: !isShareMode && !!memo?.parent,
|
||||
});
|
||||
|
||||
const {
|
||||
|
|
@ -58,7 +58,7 @@ const MemoDetail = () => {
|
|||
hasNextPage: hasNextComments,
|
||||
isFetchingNextPage: isFetchingNextComments,
|
||||
} = useInfiniteMemoComments(memoName, {
|
||||
enabled: !!memo,
|
||||
enabled: !isShareMode && !!memo,
|
||||
});
|
||||
|
||||
// Scroll to the hash target once it's in the DOM. The effect re-runs as the memo loads (footnote
|
||||
|
|
@ -80,8 +80,8 @@ const MemoDetail = () => {
|
|||
}
|
||||
}
|
||||
|
||||
// Start the memo and comment requests as soon as routing is unlocked, but do
|
||||
// not expose content before tag-blur and instance display settings settle.
|
||||
// Start the permitted requests as soon as routing is unlocked, but do not
|
||||
// expose content before tag-blur and instance display settings settle.
|
||||
if (isLoading || !memo || !authInitialized || !instanceInitialized) {
|
||||
return null;
|
||||
}
|
||||
|
|
@ -105,7 +105,7 @@ const MemoDetail = () => {
|
|||
<MentionResolutionProvider contents={mentionResolutionContents} userNames={userResolutionNames}>
|
||||
<div className={cn("w-full flex flex-row justify-start items-start px-4 sm:px-6 gap-6")}>
|
||||
<div className={cn("w-full md:w-[calc(100%-16.5rem)]")}>
|
||||
{parentMemo && (
|
||||
{!isShareMode && parentMemo && (
|
||||
<div className="w-auto inline-block mb-2">
|
||||
<Link
|
||||
className="px-3 py-1 border border-border rounded-lg max-w-xs w-auto text-sm flex flex-row justify-start items-center flex-nowrap text-muted-foreground hover:shadow hover:opacity-80"
|
||||
|
|
@ -129,14 +129,16 @@ const MemoDetail = () => {
|
|||
showPinned
|
||||
onShareImageDialogOpenChange={setShareImageDialogOpen}
|
||||
/>
|
||||
<MemoCommentSection
|
||||
memo={displayMemo}
|
||||
comments={comments}
|
||||
parentPage={locationState?.from}
|
||||
hasMoreComments={hasNextComments}
|
||||
isFetchingMoreComments={isFetchingNextComments}
|
||||
onLoadMoreComments={fetchNextComments}
|
||||
/>
|
||||
{!isShareMode && (
|
||||
<MemoCommentSection
|
||||
memo={displayMemo}
|
||||
comments={comments}
|
||||
parentPage={locationState?.from}
|
||||
hasMoreComments={hasNextComments}
|
||||
isFetchingMoreComments={isFetchingNextComments}
|
||||
onLoadMoreComments={fetchNextComments}
|
||||
/>
|
||||
)}
|
||||
</div>
|
||||
{md && (
|
||||
<div className="sticky top-0 left-0 shrink-0 -mt-6 w-60 h-full">
|
||||
|
|
|
|||
Loading…
Reference in a new issue