package interceptor
import (
"context"
"fmt"
"regexp"
"sync"
"time"
"github.com/google/trillian"
"github.com/google/trillian/monitoring"
"github.com/google/trillian/quota"
"github.com/google/trillian/quota/etcd/quotapb"
"github.com/google/trillian/server/errors"
"github.com/google/trillian/storage"
"github.com/google/trillian/trees"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
"k8s.io/klog/v2"
)
const (
badInfoReason = "bad_info"
badTreeReason = "bad_tree"
insufficientTokensReason = "insufficient_tokens"
getTreeStage = "get_tree"
getTokensStage = "get_tokens"
traceSpanRoot = "/trillian/server/int"
)
var (
PutTokensTimeout = 5 * time.Second
requestCounter monitoring.Counter
requestDeniedCounter monitoring.Counter
contextErrCounter monitoring.Counter
metricsOnce sync.Once
enabledServices = map[string]bool{
"trillian.TrillianLog": true,
"trillian.TrillianAdmin": true,
"TrillianLog": true,
"TrillianAdmin": true,
}
)
type RequestProcessor interface {
Before(ctx context.Context, req interface{}, method string) (context.Context, error)
After(ctx context.Context, resp interface{}, method string, handlerErr error)
}
type TrillianInterceptor struct {
admin storage.AdminStorage
qm quota.Manager
quotaDryRun bool
}
func New(admin storage.AdminStorage, qm quota.Manager, quotaDryRun bool, mf monitoring.MetricFactory) *TrillianInterceptor {
metricsOnce.Do(func() { initMetrics(mf) })
return &TrillianInterceptor{
admin: admin,
qm: qm,
quotaDryRun: quotaDryRun,
}
}
func initMetrics(mf monitoring.MetricFactory) {
if mf == nil {
mf = monitoring.InertMetricFactory{}
}
quota.InitMetrics(mf)
requestCounter = mf.NewCounter(
"interceptor_request_count",
"Total number of intercepted requests",
monitoring.TreeIDLabel)
requestDeniedCounter = mf.NewCounter(
"interceptor_request_denied_count",
"Number of requests by denied, labeled according to the reason for denial",
"reason", monitoring.TreeIDLabel, "quota_user")
contextErrCounter = mf.NewCounter(
"interceptor_context_err_counter",
"Total number of times request context has been cancelled or deadline exceeded by stage",
"stage")
}
func incRequestDeniedCounter(reason string, treeID int64, quotaUser string) {
requestDeniedCounter.Inc(reason, fmt.Sprint(treeID), quotaUser)
}
func (i *TrillianInterceptor) UnaryInterceptor(ctx context.Context, req interface{}, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (interface{}, error) {
rp := i.NewProcessor()
var err error
ctx, err = rp.Before(ctx, req, info.FullMethod)
if err != nil {
return nil, err
}
resp, err := handler(ctx, req)
rp.After(ctx, resp, info.FullMethod, err)
return resp, err
}
func (i *TrillianInterceptor) NewProcessor() RequestProcessor {
return &trillianProcessor{parent: i}
}
type trillianProcessor struct {
parent *TrillianInterceptor
info *rpcInfo
}
func (tp *trillianProcessor) Before(ctx context.Context, req interface{}, method string) (context.Context, error) {
if !enabledServices[serviceName(method)] {
return ctx, nil
}
innerCtx, spanEnd := spanFor(ctx, "Before")
defer spanEnd()
info, err := newRPCInfo(req)
if err != nil {
klog.Warningf("Failed to read tree info: %v", err)
incRequestDeniedCounter(badInfoReason, 0, "")
return ctx, err
}
tp.info = info
requestCounter.Inc(fmt.Sprint(info.treeID))
if info.getTree {
tree, err := trees.GetTree(
innerCtx, tp.parent.admin, info.treeID, trees.NewGetOpts(trees.Admin, info.treeTypes...))
if err != nil {
incRequestDeniedCounter(badTreeReason, info.treeID, info.quotaUsers)
return ctx, err
}
if err := innerCtx.Err(); err != nil {
contextErrCounter.Inc(getTreeStage)
return ctx, err
}
ctx = trees.NewContext(ctx, tree)
}
if info.tokens > 0 && len(info.specs) > 0 {
err := tp.parent.qm.GetTokens(innerCtx, info.tokens, info.specs)
if err != nil {
if !tp.parent.quotaDryRun {
incRequestDeniedCounter(insufficientTokensReason, info.treeID, info.quotaUsers)
return ctx, status.Errorf(codes.ResourceExhausted, "quota exhausted: %v", err)
}
klog.Warningf("(quotaDryRun) Request %+v not denied due to dry run mode: %v", req, err)
}
quota.Metrics.IncAcquired(info.tokens, info.specs, err == nil)
if err = innerCtx.Err(); err != nil {
contextErrCounter.Inc(getTokensStage)
return ctx, err
}
}
return ctx, nil
}
func (tp *trillianProcessor) After(ctx context.Context, resp interface{}, method string, handlerErr error) {
if !enabledServices[serviceName(method)] {
return
}
_, spanEnd := spanFor(ctx, "After")
defer spanEnd()
switch {
case tp.info == nil:
klog.Warningf("After called with nil rpcInfo, resp = [%+v], handlerErr = [%v]", resp, handlerErr)
return
case tp.info.tokens == 0:
return
}
refunds := make([]quota.Spec, 0)
for _, s := range tp.info.specs {
if s.Refundable {
refunds = append(refunds, s)
}
}
if len(refunds) == 0 {
return
}
tokens := 0
if handlerErr != nil {
tokens = tp.info.tokens
} else {
switch resp := resp.(type) {
case *trillian.QueueLeafResponse:
if !isLeafOK(resp.GetQueuedLeaf()) {
tokens = 1
}
case *trillian.AddSequencedLeavesResponse:
for _, leaf := range resp.GetResults() {
if !isLeafOK(leaf) {
tokens++
}
}
}
}
if tokens > 0 {
go func() {
ctx, spanEnd := spanFor(context.Background(), "After.PutTokens")
defer spanEnd()
ctx, cancel := context.WithTimeout(ctx, PutTokensTimeout)
defer cancel()
err := tp.parent.qm.PutTokens(ctx, tokens, refunds)
if err != nil {
klog.Warningf("Failed to replenish %v tokens: %v", tokens, err)
}
quota.Metrics.IncReturned(tokens, refunds, err == nil)
}()
}
}
func isLeafOK(leaf *trillian.QueuedLogLeaf) bool {
return leaf == nil || leaf.Status == nil || leaf.Status.Code == int32(codes.OK)
}
var (
fullyQualifiedRE = regexp.MustCompile(`^/([\w.]+)/(\w+)$`)
unqualifiedRE = regexp.MustCompile(`^/(\w+)\.(\w+)$`)
)
func serviceName(fullMethod string) string {
if matches := fullyQualifiedRE.FindStringSubmatch(fullMethod); len(matches) == 3 {
return matches[1]
}
if matches := unqualifiedRE.FindStringSubmatch(fullMethod); len(matches) == 3 {
return matches[1]
}
return ""
}
type rpcInfo struct {
getTree bool
readonly bool
treeID int64
treeTypes []trillian.TreeType
specs []quota.Spec
tokens int
quotaUsers string
}
type chargable interface {
GetChargeTo() *trillian.ChargeTo
}
func chargedUsers(req interface{}) []string {
c, ok := req.(chargable)
if !ok {
return nil
}
chargeTo := c.GetChargeTo()
if chargeTo == nil {
return nil
}
return chargeTo.User
}
func newRPCInfoForRequest(req interface{}) (*rpcInfo, error) {
info := &rpcInfo{
getTree: true,
readonly: true,
treeTypes: nil,
tokens: 0,
}
switch req := req.(type) {
case
*quotapb.CreateConfigRequest,
*quotapb.DeleteConfigRequest,
*quotapb.GetConfigRequest,
*quotapb.ListConfigsRequest,
*quotapb.UpdateConfigRequest:
info.getTree = false
info.readonly = false
case *trillian.CreateTreeRequest:
info.getTree = false
info.readonly = false
case *trillian.ListTreesRequest:
info.getTree = false
case *trillian.GetTreeRequest:
info.getTree = false
case *trillian.DeleteTreeRequest,
*trillian.UndeleteTreeRequest,
*trillian.UpdateTreeRequest:
info.getTree = false
info.readonly = false
case *trillian.GetConsistencyProofRequest,
*trillian.GetEntryAndProofRequest,
*trillian.GetInclusionProofByHashRequest,
*trillian.GetInclusionProofRequest,
*trillian.GetLatestSignedLogRootRequest:
info.treeTypes = []trillian.TreeType{trillian.TreeType_LOG, trillian.TreeType_PREORDERED_LOG}
info.tokens = 1
case *trillian.GetLeavesByRangeRequest:
info.treeTypes = []trillian.TreeType{trillian.TreeType_LOG, trillian.TreeType_PREORDERED_LOG}
info.tokens = 1
if c := req.GetCount(); c > 1 {
info.tokens = int(c)
}
case *trillian.QueueLeafRequest:
info.readonly = false
info.treeTypes = []trillian.TreeType{trillian.TreeType_LOG}
info.tokens = 1
case *trillian.AddSequencedLeavesRequest:
info.readonly = false
info.treeTypes = []trillian.TreeType{trillian.TreeType_PREORDERED_LOG}
info.tokens = len(req.GetLeaves())
case *trillian.InitLogRequest:
info.readonly = false
info.treeTypes = []trillian.TreeType{trillian.TreeType_LOG, trillian.TreeType_PREORDERED_LOG}
info.tokens = 1
default:
return nil, status.Errorf(codes.Internal, "newRPCInfo: unmapped request type: %T", req)
}
return info, nil
}
func newRPCInfo(req interface{}) (*rpcInfo, error) {
info, err := newRPCInfoForRequest(req)
if err != nil {
return nil, err
}
if info.getTree || info.tokens > 0 {
switch req := req.(type) {
case logIDRequest:
info.treeID = req.GetLogId()
case treeIDRequest:
info.treeID = req.GetTreeId()
case treeRequest:
info.treeID = req.GetTree().GetTreeId()
default:
return nil, status.Errorf(codes.Internal, "cannot retrieve treeID from request: %T", req)
}
}
if info.tokens > 0 {
kind := quota.Write
if info.readonly {
kind = quota.Read
}
for _, user := range chargedUsers(req) {
info.specs = append(info.specs, quota.Spec{Group: quota.User, Kind: kind, User: user})
if len(info.quotaUsers) > 0 {
info.quotaUsers += "+"
}
info.quotaUsers += user
}
info.specs = append(info.specs, []quota.Spec{
{Group: quota.Tree, Kind: kind, TreeID: info.treeID},
{Group: quota.Global, Kind: kind, Refundable: true},
}...)
}
return info, nil
}
type logIDRequest interface {
GetLogId() int64
}
type treeIDRequest interface {
GetTreeId() int64
}
type treeRequest interface {
GetTree() *trillian.Tree
}
func ErrorWrapper(ctx context.Context, req interface{}, _ *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (interface{}, error) {
ctx, spanEnd := spanFor(ctx, "ErrorWrapper")
defer spanEnd()
rsp, err := handler(ctx, req)
return rsp, errors.WrapError(err)
}
func spanFor(ctx context.Context, name string) (context.Context, func()) {
return monitoring.StartSpan(ctx, fmt.Sprintf("%s.%s", traceSpanRoot, name))
}