feat(storage): policy-level content audit on share creation (upstream #2178 item 9)
Generated with [Devin](https://devin.ai) Co-Authored-By: Devin <158243242+devin-ai-integration[bot]@users.noreply.github.com>pull/3593/head
parent
28d927cf7b
commit
27c8b05f1a
@ -0,0 +1,125 @@
|
|||||||
|
package manager
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/ent"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/inventory/types"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/activity"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/filemanager/fs"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/request"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/serializer"
|
||||||
|
)
|
||||||
|
|
||||||
|
// auditTimeout bounds a single audit request so a stuck endpoint cannot hang
|
||||||
|
// share creation.
|
||||||
|
const auditTimeout = 30 * time.Second
|
||||||
|
|
||||||
|
type auditResponse struct {
|
||||||
|
Flagged bool `json:"flagged"`
|
||||||
|
Reason string `json:"reason"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// auditShareEntities submits every auditable entity covered by a share to the
|
||||||
|
// owning policy's audit endpoint (PolicySetting.AuditEndpoint, 鉴黄). Only
|
||||||
|
// image entities within AuditMaxSize are sent; a flagged result aborts share
|
||||||
|
// creation. Endpoint failures are fail-closed — an unavailable audit service
|
||||||
|
// must not silently bypass the check.
|
||||||
|
func (l *manager) auditShareEntities(ctx context.Context, files []fs.File) error {
|
||||||
|
policyClient := l.dep.StoragePolicyClient()
|
||||||
|
mime := l.dep.MimeDetector(ctx)
|
||||||
|
|
||||||
|
seen := map[int]struct{}{}
|
||||||
|
policies := map[int]*ent.StoragePolicy{}
|
||||||
|
for _, file := range files {
|
||||||
|
// The extension lives on the file name; entity sources are blob keys.
|
||||||
|
fileMime := mime.TypeByName(file.Name())
|
||||||
|
if !strings.HasPrefix(fileMime, "image/") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
for _, e := range file.Entities() {
|
||||||
|
if _, ok := seen[e.ID()]; ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
seen[e.ID()] = struct{}{}
|
||||||
|
|
||||||
|
if e.Type() != types.EntityTypeVersion || e.Size() == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
policy, ok := policies[e.PolicyID()]
|
||||||
|
if !ok {
|
||||||
|
p, err := policyClient.GetPolicyByID(ctx, e.PolicyID())
|
||||||
|
if err != nil {
|
||||||
|
return serializer.NewError(serializer.CodeContentAuditFailed, "failed to load entity storage policy", err)
|
||||||
|
}
|
||||||
|
policy = p
|
||||||
|
policies[e.PolicyID()] = p
|
||||||
|
}
|
||||||
|
|
||||||
|
endpoint := policy.Settings.AuditEndpoint
|
||||||
|
if endpoint == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if policy.Settings.AuditMaxSize > 0 && e.Size() > policy.Settings.AuditMaxSize {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
src, err := l.GetEntitySource(ctx, e.ID(), fs.WithEntity(e))
|
||||||
|
if err != nil {
|
||||||
|
return serializer.NewError(serializer.CodeContentAuditFailed, "failed to open entity for audit", err)
|
||||||
|
}
|
||||||
|
flagged, reason, err := callAuditEndpoint(l.dep.RequestClient(request.WithContext(ctx), request.WithTimeout(auditTimeout)),
|
||||||
|
endpoint, file.Name(), e.Size(), fileMime, src)
|
||||||
|
src.Close()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if flagged {
|
||||||
|
activity.Record(ctx, l.settings, l.dep.ActivityClient(), types.EventContentAuditBlocked,
|
||||||
|
activity.File(file.ID()), activity.Extra(map[string]any{
|
||||||
|
"entity_id": e.ID(),
|
||||||
|
"reason": reason,
|
||||||
|
}))
|
||||||
|
return serializer.NewError(serializer.CodeContentAuditRejected, "content rejected by audit", nil)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// callAuditEndpoint POSTs the entity body to the audit endpoint and returns the
|
||||||
|
// flagged verdict. Endpoint must be plain http(s); the endpoint answers
|
||||||
|
// {"flagged": bool, "reason": string}.
|
||||||
|
func callAuditEndpoint(client request.Client, endpoint, name string, size int64, mimeType string, body io.Reader) (bool, string, error) {
|
||||||
|
u, err := url.Parse(endpoint)
|
||||||
|
if err != nil || (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" {
|
||||||
|
return false, "", serializer.NewError(serializer.CodeContentAuditFailed, "invalid audit endpoint", nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
resp := client.Request(http.MethodPost, endpoint, body,
|
||||||
|
request.WithContentLength(size),
|
||||||
|
request.WithHeader(http.Header{
|
||||||
|
"Content-Type": {mimeType},
|
||||||
|
"X-Entity-Name": {name},
|
||||||
|
"X-Entity-Size": {strconv.FormatInt(size, 10)},
|
||||||
|
}),
|
||||||
|
)
|
||||||
|
raw, err := resp.CheckHTTPResponse(http.StatusOK).GetResponse()
|
||||||
|
if err != nil {
|
||||||
|
return false, "", serializer.NewError(serializer.CodeContentAuditFailed, "audit endpoint error", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var result auditResponse
|
||||||
|
if err := json.Unmarshal([]byte(raw), &result); err != nil {
|
||||||
|
return false, "", serializer.NewError(serializer.CodeContentAuditFailed, "invalid audit response", err)
|
||||||
|
}
|
||||||
|
return result.Flagged, result.Reason, nil
|
||||||
|
}
|
||||||
@ -0,0 +1,82 @@
|
|||||||
|
package manager
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/request"
|
||||||
|
"github.com/cloudreve/Cloudreve/v4/pkg/serializer"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newAuditTestClient(t *testing.T) request.Client {
|
||||||
|
return request.NewClient(nil, request.WithTimeout(0))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCallAuditEndpoint(t *testing.T) {
|
||||||
|
t.Run("flagged", func(t *testing.T) {
|
||||||
|
var gotBody []byte
|
||||||
|
var gotName, gotSize, gotMime string
|
||||||
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
gotBody, _ = io.ReadAll(r.Body)
|
||||||
|
gotName = r.Header.Get("X-Entity-Name")
|
||||||
|
gotSize = r.Header.Get("X-Entity-Size")
|
||||||
|
gotMime = r.Header.Get("Content-Type")
|
||||||
|
w.Write([]byte(`{"flagged":true,"reason":"nsfw"}`))
|
||||||
|
}))
|
||||||
|
t.Cleanup(srv.Close)
|
||||||
|
|
||||||
|
flagged, reason, err := callAuditEndpoint(newAuditTestClient(t), srv.URL, "pic.png", 4, "image/png", strings.NewReader("data"))
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.True(t, flagged)
|
||||||
|
require.Equal(t, "nsfw", reason)
|
||||||
|
require.Equal(t, "data", string(gotBody))
|
||||||
|
require.Equal(t, "pic.png", gotName)
|
||||||
|
require.Equal(t, "4", gotSize)
|
||||||
|
require.Equal(t, "image/png", gotMime)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("clean", func(t *testing.T) {
|
||||||
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.Write([]byte(`{"flagged":false}`))
|
||||||
|
}))
|
||||||
|
t.Cleanup(srv.Close)
|
||||||
|
|
||||||
|
flagged, _, err := callAuditEndpoint(newAuditTestClient(t), srv.URL, "pic.png", 4, "image/png", strings.NewReader("data"))
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.False(t, flagged)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("fail closed on endpoint error", func(t *testing.T) {
|
||||||
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.WriteHeader(http.StatusInternalServerError)
|
||||||
|
}))
|
||||||
|
t.Cleanup(srv.Close)
|
||||||
|
|
||||||
|
_, _, err := callAuditEndpoint(newAuditTestClient(t), srv.URL, "pic.png", 4, "image/png", strings.NewReader("data"))
|
||||||
|
require.Error(t, err)
|
||||||
|
require.Equal(t, serializer.CodeContentAuditFailed, err.(serializer.AppError).Code)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("fail closed on bad json", func(t *testing.T) {
|
||||||
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.Write([]byte(`not json`))
|
||||||
|
}))
|
||||||
|
t.Cleanup(srv.Close)
|
||||||
|
|
||||||
|
_, _, err := callAuditEndpoint(newAuditTestClient(t), srv.URL, "pic.png", 4, "image/png", strings.NewReader("data"))
|
||||||
|
require.Error(t, err)
|
||||||
|
require.Equal(t, serializer.CodeContentAuditFailed, err.(serializer.AppError).Code)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("rejects non-http schemes", func(t *testing.T) {
|
||||||
|
for _, endpoint := range []string{"file:///etc/passwd", "ftp://x/", "://bad", ""} {
|
||||||
|
_, _, err := callAuditEndpoint(newAuditTestClient(t), endpoint, "pic.png", 4, "image/png", strings.NewReader("data"))
|
||||||
|
require.Error(t, err, endpoint)
|
||||||
|
require.Equal(t, serializer.CodeContentAuditFailed, err.(serializer.AppError).Code)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
Loading…
Reference in new issue