Update module google.golang.org/grpc to v1.84.0

Signed-off-by: renovate[bot] <29139614+renovate[bot]@users.noreply.github.com>
This commit is contained in:
renovate[bot] 2026-09-18 05:32:34 +00:00 • committed by GitHub
parent 6ba4ab29fc
commit fca1cae7a6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
25 changed files with 781 additions and 739 deletions

2
go.mod
View file

@ -73,7 +73,7 @@ require (
golang.org/x/sync v0.23.0
golang.org/x/sys v0.48.0
golang.org/x/term v0.46.0
google.golang.org/grpc v1.83.2
google.golang.org/grpc v1.84.0
google.golang.org/protobuf v1.36.12
gopkg.in/inf.v0 v0.9.1
gopkg.in/yaml.v3 v3.0.1

4
go.sum
View file

@ -554,8 +554,8 @@ google.golang.org/genproto/googleapis/api v0.0.0-20260825221802-da73d73af1c5 h1:
google.golang.org/genproto/googleapis/api v0.0.0-20260825221802-da73d73af1c5/go.mod h1:3LhxRw4YYkf+ylAfgaY9JlVLFKhokkCV8duhLLe7+t0=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260825221802-da73d73af1c5 h1:1VUiZAXyC+zmiFYi+WLtBzr68Cj8wOofHjjrA/kkizc=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260825221802-da73d73af1c5/go.mod h1:DjtHYE8FKJLivXcBEjGwndXfIC23G0VpXiXKqG179uA=
google.golang.org/grpc v1.83.2 h1:EManeRomTObA0BU7I8vXgg/78uE5MJ9M8B39EX2WscU=
google.golang.org/grpc v1.83.2/go.mod h1:YPI1hK3kDked6iHvgX3tR0y+nX/qpMFKhPgFsokw1S8=
google.golang.org/grpc v1.84.0 h1:soMyaPJ8pAak5PIQ0DGBUir0XRo2fRoMqhNWMLlLxO0=
google.golang.org/grpc v1.84.0/go.mod h1:ljCht0DrxQrXBDRTZp52Qxh3Ffk8CdYm2sj4O2QN2C0=
google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc=
google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=

View file

@ -39,27 +39,14 @@ import (
var randIntN = rand.IntN
// ChildState is the balancer state of a child along with the endpoint which
// identifies the child balancer.
// ChildState is the state of a child balancer.
type ChildState struct {
Endpoint resolver.Endpoint
State balancer.State
// Balancer exposes only the ExitIdler interface of the child LB policy.
// Other methods of the child policy are called only by endpointsharding.
Balancer ExitIdler
Endpoint resolver.Endpoint // Endpoint of the child balancer.
State balancer.State // State of the child balancer.
ExitIdle func() // Function to exit the child balancer from IDLE state.
}
// ExitIdler provides access to only the ExitIdle method of the child balancer.
type ExitIdler interface {
// ExitIdle instructs the LB policy to reconnect to backends / exit the
// IDLE state, if appropriate and possible. Note that SubConns that enter
// the IDLE state will not reconnect until SubConn.Connect is called.
ExitIdle()
}
// Options are the options to configure the behaviour of the
// endpointsharding balancer.
// Options configure the behaviour of the endpointsharding balancer.
type Options struct {
// DisableAutoReconnect allows the balancer to keep child balancer in the
// IDLE state until they are explicitly triggered to exit using the
@ -77,14 +64,13 @@ type ChildBuilderFunc func(cc balancer.ClientConn, opts balancer.BuildOptions) b
// policies each owning a single endpoint. The endpointsharding balancer
// forwards the LoadBalancingConfig in ClientConn state updates to its children.
func NewBalancer(cc balancer.ClientConn, opts balancer.BuildOptions, childBuilder ChildBuilderFunc, esOpts Options) balancer.Balancer {
es := &endpointSharding{
return &endpointSharding{
cc: cc,
bOpts: opts,
esOpts: esOpts,
childBuilder: childBuilder,
endpoints: resolver.NewEndpointMap[*endpointState](),
}
es.children.Store(resolver.NewEndpointMap[*balancerWrapper]())
return es
}
// endpointSharding is a balancer that wraps child balancers. It creates a child
@ -96,124 +82,129 @@ type endpointSharding struct {
esOpts Options
childBuilder ChildBuilderFunc
// childMu synchronizes calls to any single child. It must be held for all
// calls into a child. To avoid deadlocks, do not acquire childMu while
// holding mu.
childMu sync.Mutex
children atomic.Pointer[resolver.EndpointMap[*balancerWrapper]]
// inhibitChildUpdates is set during UpdateClientConnState/ResolverError
// calls (calls to children will each produce an update, only want one
// update).
inhibitChildUpdates atomic.Bool
// mu synchronizes access to the state stored in balancerWrappers in the
// children field. mu must not be held during calls into a child since
// synchronous calls back from the child may require taking mu, causing a
// deadlock. To avoid deadlocks, do not acquire childMu while holding mu.
mu sync.Mutex
// mu guards access to the below fields and guarantees mutual exclusion
// between top-down methods (like UpdateClientConnState, ResolverError etc,
// which are already serialized) and bottom-up methods (like UpdateState)
// that can be called concurrently.
//
// Lock ordering & deadlock prevention:
// Child callbacks (like UpdateState) acquire this mutex to push state updates
// to the parent. This establishes a strict lock ordering:
// [endpointState.childMu] -> [endpointSharding.mu]
//
// To prevent deadlocks, we must never invert this order. Therefore, we must
// never call any methods on a child balancer while holding this mutex. If we
// did, and that child balancer invoked UpdateState synchronously, it would
// attempt to re-acquire this mutex, causing a deadlock.
//
// Concurrency:
// Top-down operations need to iterate over all children. To ensure a
// single, clean aggregated update at the end of such operations, we inhibit
// intermediate updates from children. We grab this mutex briefly to set
// `inhibitChildUpdates = true` and immediately release it. This allows us
// to perform the potentially slow, top-down child updates without holding
// any parent locks. Once finished, we grab the mutex again to unset the
// flag and push the final aggregated state. An update from a child during
// this time will *only* update the child state, and will not access the
// endpoints map or push an aggregated state to the parent.
mu sync.Mutex
endpoints *resolver.EndpointMap[*endpointState]
inhibitChildUpdates bool
}
// rotateEndpoints returns a slice of all the input endpoints rotated a random
// amount.
func rotateEndpoints(es []resolver.Endpoint) []resolver.Endpoint {
les := len(es)
if les == 0 {
n := len(es)
if n == 0 {
return es
}
r := randIntN(les)
r := randIntN(n)
// Make a copy to avoid mutating data beyond the end of es.
ret := make([]resolver.Endpoint, les)
ret := make([]resolver.Endpoint, n)
copy(ret, es[r:])
copy(ret[les-r:], es[:r])
copy(ret[n-r:], es[:r])
return ret
}
// UpdateClientConnState creates a child for new endpoints and deletes children
// for endpoints that are no longer present. It also updates all the children,
// and sends a single synchronous update of the childrens' aggregated state at
// the end of the UpdateClientConnState operation. If any endpoint has no
// addresses it will ignore that endpoint. Otherwise, returns first error found
// from a child, but fully processes the new update.
// the end of the UpdateClientConnState operation.
//
// Returns the first error found from a child, but fully processes the update.
func (es *endpointSharding) UpdateClientConnState(state balancer.ClientConnState) error {
es.childMu.Lock()
defer es.childMu.Unlock()
es.inhibitUpdatesFromChildren()
es.inhibitChildUpdates.Store(true)
defer func() {
es.inhibitChildUpdates.Store(false)
es.updateState()
}()
var ret error
children := es.children.Load()
newChildren := resolver.NewEndpointMap[*balancerWrapper]()
// Update/Create new children.
// Update/create child balancers for each endpoint in the update. Note that we
// don't hold the mutex here, but this is fine because inhibitChildUpdates is
// true, and therefore UpdateState will not access es.endpoints.
var retErr error
newEndpoints := resolver.NewEndpointMap[*endpointState]()
for _, endpoint := range rotateEndpoints(state.ResolverState.Endpoints) {
if _, ok := newChildren.Get(endpoint); ok {
// Endpoint child was already created, continue to avoid duplicate
// update.
if _, ok := newEndpoints.Get(endpoint); ok {
// Skip duplicate endpoints.
continue
}
childBalancer, ok := children.Get(endpoint)
epState, ok := es.endpoints.Get(endpoint)
if ok {
// Endpoint attributes may have changed, update the stored endpoint.
es.mu.Lock()
childBalancer.childState.Endpoint = endpoint
es.mu.Unlock()
// Endpoint child already exists, update the stored endpoint.
epState.endpoint = endpoint
} else {
childBalancer = &balancerWrapper{
childState: ChildState{Endpoint: endpoint},
ClientConn: es.cc,
es: es,
// Endpoint child does not exist, create a new one.
epState = &endpointState{
ClientConn: es.cc,
parent: es,
endpoint: endpoint,
disableAutoReconnect: es.esOpts.DisableAutoReconnect,
}
childBalancer.childState.Balancer = childBalancer
childBalancer.child = es.childBuilder(childBalancer, es.bOpts)
epState.childLB = es.childBuilder(epState, es.bOpts)
}
newChildren.Set(endpoint, childBalancer)
if err := childBalancer.updateClientConnStateLocked(balancer.ClientConnState{
// Update the endpoint state for the endpoint.
newEndpoints.Set(endpoint, epState)
if err := epState.updateClientConnState(balancer.ClientConnState{
BalancerConfig: state.BalancerConfig,
ResolverState: resolver.State{
Endpoints: []resolver.Endpoint{endpoint},
Attributes: state.ResolverState.Attributes,
},
}); err != nil && ret == nil {
// Return first error found, and always commit full processing of
// updating children. If desired to process more specific errors
// across all endpoints, caller should make these specific
// validations, this is a current limitation for simplicity sake.
ret = err
}); err != nil && retErr == nil {
// Keep the first error found from any child.
retErr = err
}
}
// Delete old children that are no longer present.
for e, child := range children.All() {
if _, ok := newChildren.Get(e); !ok {
child.closeLocked()
for e, child := range es.endpoints.All() {
if _, ok := newEndpoints.Get(e); !ok {
child.close()
}
}
es.children.Store(newChildren)
if newChildren.Len() == 0 {
return balancer.ErrBadResolverState
if newEndpoints.Len() == 0 {
retErr = balancer.ErrBadResolverState
}
return ret
es.mu.Lock()
es.endpoints = newEndpoints
es.inhibitChildUpdates = false
es.updateStateLocked()
es.mu.Unlock()
return retErr
}
// ResolverError forwards the resolver error to all of the endpointSharding's
// children and sends a single synchronous update of the childStates at the end
// of the ResolverError operation.
func (es *endpointSharding) ResolverError(err error) {
es.childMu.Lock()
defer es.childMu.Unlock()
es.inhibitChildUpdates.Store(true)
defer func() {
es.inhibitChildUpdates.Store(false)
es.updateState()
}()
children := es.children.Load()
for _, child := range children.All() {
child.resolverErrorLocked(err)
es.inhibitUpdatesFromChildren()
for _, child := range es.endpoints.All() {
child.resolverError(err)
}
es.allowUpdatesFromChildren()
}
func (es *endpointSharding) UpdateSubConnState(balancer.SubConn, balancer.SubConnState) {
@ -221,41 +212,49 @@ func (es *endpointSharding) UpdateSubConnState(balancer.SubConn, balancer.SubCon
}
func (es *endpointSharding) Close() {
es.childMu.Lock()
defer es.childMu.Unlock()
children := es.children.Load()
for _, child := range children.All() {
child.closeLocked()
es.inhibitUpdatesFromChildren()
for _, child := range es.endpoints.All() {
child.close()
}
}
func (es *endpointSharding) ExitIdle() {
es.childMu.Lock()
defer es.childMu.Unlock()
for _, bw := range es.children.Load().All() {
if !bw.isClosed {
bw.child.ExitIdle()
}
es.inhibitUpdatesFromChildren()
for _, child := range es.endpoints.All() {
child.exitIdle()
}
es.allowUpdatesFromChildren()
}
// updateState updates this component's state. It sends the aggregated state,
// and a picker with round robin behavior with all the child states present if
// needed.
func (es *endpointSharding) updateState() {
if es.inhibitChildUpdates.Load() {
return
}
func (es *endpointSharding) inhibitUpdatesFromChildren() {
es.mu.Lock()
es.inhibitChildUpdates = true
es.mu.Unlock()
}
func (es *endpointSharding) allowUpdatesFromChildren() {
es.mu.Lock()
es.inhibitChildUpdates = false
es.updateStateLocked()
es.mu.Unlock()
}
// updateStateLocked updates this component's state. It sends the aggregated
// state, and a picker with round robin behavior with all the child states
// present if needed. This method must only be called when inhibitChildUpdates
// is false.
//
// Caller must hold es.mu.
func (es *endpointSharding) updateStateLocked() {
var readyPickers, connectingPickers, idlePickers, transientFailurePickers []balancer.Picker
es.mu.Lock()
defer es.mu.Unlock()
children := es.children.Load()
childStates := make([]ChildState, 0, children.Len())
for _, child := range children.All() {
childState := child.childState
childStates := make([]ChildState, 0, es.endpoints.Len())
for _, epState := range es.endpoints.All() {
childState := ChildState{
Endpoint: epState.endpoint,
State: epState.state,
ExitIdle: func() { go epState.exitIdle() },
}
childStates = append(childStates, childState)
childPicker := childState.State.Picker
switch childState.State.ConnectivityState {
@ -292,14 +291,14 @@ func (es *endpointSharding) updateState() {
aggState = connectivity.TransientFailure
pickers = []balancer.Picker{base.NewErrPicker(errors.New("no children to pick from"))}
} // No children (resolver error before valid update).
p := &pickerWithChildStates{
pickers: pickers,
childStates: childStates,
next: uint32(randIntN(len(pickers))),
}
es.cc.UpdateState(balancer.State{
ConnectivityState: aggState,
Picker: p,
Picker: &pickerWithChildStates{
pickers: pickers,
childStates: childStates,
next: uint32(randIntN(len(pickers))),
},
})
}
@ -328,61 +327,59 @@ func ChildStatesFromPicker(picker balancer.Picker) []ChildState {
return p.childStates
}
// balancerWrapper is a wrapper of a balancer. It ID's a child balancer by
// endpoint, and persists recent child balancer state.
type balancerWrapper struct {
// The following fields are initialized at build time and read-only after
// that and therefore do not need to be guarded by a mutex.
// endpointState is the internal state maintained for each endpoint.
type endpointState struct {
balancer.ClientConn // Embedded to intercept UpdateState
// child contains the wrapped balancer. Access its methods only through
// methods on balancerWrapper to ensure proper synchronization
child balancer.Balancer
balancer.ClientConn // embed to intercept UpdateState, doesn't deal with SubConns
parent *endpointSharding // Parent endpointsharding balancer.
endpoint resolver.Endpoint // Endpoint of the child balancer.
state balancer.State // State of the child balancer.
disableAutoReconnect bool // Whether to disable auto reconnect for this child.
es *endpointSharding
// Access to the following fields is guarded by es.mu.
childState ChildState
isClosed bool
childMu sync.Mutex // Guarantees mutual exclusion for Balancer API calls to the child balancer.
childLB balancer.Balancer // Child balancer.
closed bool // Tracks closure of the child balancer to ensure ExitIdle is not called after Close().
}
func (bw *balancerWrapper) UpdateState(state balancer.State) {
bw.es.mu.Lock()
bw.childState.State = state
bw.es.mu.Unlock()
if state.ConnectivityState == connectivity.Idle && !bw.es.esOpts.DisableAutoReconnect {
bw.ExitIdle()
func (es *endpointState) UpdateState(state balancer.State) {
es.parent.mu.Lock()
es.state = state
if !es.parent.inhibitChildUpdates {
es.parent.updateStateLocked()
}
es.parent.mu.Unlock()
if state.ConnectivityState == connectivity.Idle && !es.disableAutoReconnect {
go es.exitIdle()
}
bw.es.updateState()
}
// ExitIdle pings an IDLE child balancer to exit idle in a new goroutine to
// avoid deadlocks due to synchronous balancer state updates.
func (bw *balancerWrapper) ExitIdle() {
go func() {
bw.es.childMu.Lock()
if !bw.isClosed {
bw.child.ExitIdle()
}
bw.es.childMu.Unlock()
}()
func (es *endpointState) updateClientConnState(state balancer.ClientConnState) error {
es.childMu.Lock()
err := es.childLB.UpdateClientConnState(state)
es.childMu.Unlock()
return err
}
// updateClientConnStateLocked delivers the ClientConnState to the child
// balancer. Callers must hold the child mutex of the parent endpointsharding
// balancer.
func (bw *balancerWrapper) updateClientConnStateLocked(ccs balancer.ClientConnState) error {
return bw.child.UpdateClientConnState(ccs)
func (es *endpointState) resolverError(err error) {
es.childMu.Lock()
es.childLB.ResolverError(err)
es.childMu.Unlock()
}
// closeLocked closes the child balancer. Callers must hold the child mutext of
// the parent endpointsharding balancer.
func (bw *balancerWrapper) closeLocked() {
bw.child.Close()
bw.isClosed = true
func (es *endpointState) close() {
es.childMu.Lock()
if !es.closed {
es.closed = true
es.childLB.Close()
}
es.childMu.Unlock()
}
func (bw *balancerWrapper) resolverErrorLocked(err error) {
bw.child.ResolverError(err)
func (es *endpointState) exitIdle() {
es.childMu.Lock()
if !es.closed {
es.childLB.ExitIdle()
}
es.childMu.Unlock()
}

View file

@ -36,7 +36,6 @@ import (
"google.golang.org/grpc/balancer/pickfirst/internal"
"google.golang.org/grpc/connectivity"
"google.golang.org/grpc/experimental/balancer/weight"
expstats "google.golang.org/grpc/experimental/stats"
"google.golang.org/grpc/grpclog"
"google.golang.org/grpc/internal/envconfig"
internalgrpclog "google.golang.org/grpc/internal/grpclog"
@ -56,30 +55,7 @@ const Name = "pick_first"
// attributes to indicate whether the health listener usage is enabled.
type enableHealthListenerKeyType struct{}
var (
logger = grpclog.Component("pick-first-leaf-lb")
disconnectionsMetric = expstats.RegisterInt64Count(expstats.MetricDescriptor{
Name: "grpc.lb.pick_first.disconnections",
Description: "EXPERIMENTAL. Number of times the selected subchannel becomes disconnected.",
Unit: "{disconnection}",
Labels: []string{"grpc.target"},
Default: false,
})
connectionAttemptsSucceededMetric = expstats.RegisterInt64Count(expstats.MetricDescriptor{
Name: "grpc.lb.pick_first.connection_attempts_succeeded",
Description: "EXPERIMENTAL. Number of successful connection attempts.",
Unit: "{attempt}",
Labels: []string{"grpc.target"},
Default: false,
})
connectionAttemptsFailedMetric = expstats.RegisterInt64Count(expstats.MetricDescriptor{
Name: "grpc.lb.pick_first.connection_attempts_failed",
Description: "EXPERIMENTAL. Number of failed connection attempts.",
Unit: "{attempt}",
Labels: []string{"grpc.target"},
Default: false,
})
)
var logger = grpclog.Component("pick-first-leaf-lb")
const (
// TODO: change to pick-first when this becomes the default pick_first policy.
@ -101,11 +77,9 @@ const (
type pickfirstBuilder struct{}
func (pickfirstBuilder) Build(cc balancer.ClientConn, bo balancer.BuildOptions) balancer.Balancer {
func (pickfirstBuilder) Build(cc balancer.ClientConn, _ balancer.BuildOptions) balancer.Balancer {
b := &pickfirstBalancer{
cc: cc,
target: bo.Target.String(),
metricsRecorder: cc.MetricsRecorder(),
cc: cc,
subConns: resolver.NewAddressMapV2[*scData](),
state: connectivity.Connecting,
@ -185,10 +159,8 @@ func (b *pickfirstBalancer) newSCData(addr resolver.Address) (*scData, error) {
type pickfirstBalancer struct {
// The following fields are initialized at build time and read-only after
// that and therefore do not need to be guarded by a mutex.
logger *internalgrpclog.PrefixLogger
cc balancer.ClientConn
target string
metricsRecorder expstats.MetricsRecorder // guaranteed to be non nil
logger *internalgrpclog.PrefixLogger
cc balancer.ClientConn
// The mutex is used to ensure synchronization of updates triggered
// from the idle picker and the already serialized resolver,
@ -633,14 +605,11 @@ func (b *pickfirstBalancer) updateSubConnState(sd *scData, newState balancer.Sub
return
}
// Record a connection attempt when exiting CONNECTING.
if newState.ConnectivityState == connectivity.TransientFailure {
sd.connectionFailedInFirstPass = true
connectionAttemptsFailedMetric.Record(b.metricsRecorder, 1, b.target)
}
if newState.ConnectivityState == connectivity.Ready {
connectionAttemptsSucceededMetric.Record(b.metricsRecorder, 1, b.target)
b.shutdownRemainingLocked(sd)
if !b.addressList.seekTo(sd.addr) {
// This should not fail as we should have only one SubConn after
@ -689,15 +658,6 @@ func (b *pickfirstBalancer) updateSubConnState(sd *scData, newState balancer.Sub
// the first address when the picker is used.
b.shutdownRemainingLocked(sd)
sd.effectiveState = newState.ConnectivityState
// READY SubConn interspliced in between CONNECTING and IDLE, need to
// account for that.
if oldState == connectivity.Connecting {
// A known issue (https://github.com/grpc/grpc-go/issues/7862)
// causes a race that prevents the READY state change notification.
// This works around it.
connectionAttemptsSucceededMetric.Record(b.metricsRecorder, 1, b.target)
}
disconnectionsMetric.Record(b.metricsRecorder, 1, b.target)
b.addressList.reset()
b.updateBalancerState(balancer.State{
ConnectivityState: connectivity.Idle,

View file

@ -679,7 +679,7 @@ func (x *Message) GetData() []byte {
// A list of metadata pairs, used in the payload of client header,
// server header, and server trailer.
// Implementations may omit some entries to honor the header limits
// of GRPC_BINARY_LOG_CONFIG.
// of GRPC_BINARY_LOG_FILTER.
//
// Header keys added by gRPC are omitted. To be more specific,
// implementations will not log the following entries, and this is
@ -689,7 +689,9 @@ func (x *Message) GetData() []byte {
// or keys like 'lb-token'
// - transport specific entries, including but not limited to:
// ':path', ':authority', 'content-encoding', 'user-agent', 'te', etc
// - entries added for call credentials
// - entries added for call credentials (on client side, excludes all
// such headers; on server side, excludes only the "authorization"
// header)
//
// Implementations must always log grpc-trace-bin if it is present.
// Practically speaking it will only be visible on server side because

View file

@ -20,6 +20,10 @@ package grpc
import (
"context"
"io"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
)
// Invoke sends the RPC request on the wire and returns after response is
@ -67,8 +71,27 @@ func invoke(ctx context.Context, method string, req, reply any, cc *ClientConn,
if err != nil {
return err
}
if err := cs.SendMsg(req); err != nil {
// In case of nil or io.EOF error, call RecvMsg.
if err := cs.SendMsg(req); err != nil && err != io.EOF {
return err
}
return cs.RecvMsg(reply)
// CloseSend is needed because in some scenarios (e.g., xDS), the same
// interceptors are used to process both unary and streaming RPCs. Calling
// CloseSend signals to those interceptors that no more messages are on the
// way.
if err := cs.CloseSend(); err != nil && err != io.EOF {
return err
}
if err := cs.RecvMsg(reply); err != nil {
return err
}
// Call RecvMsg again to get the trailers.
err = cs.RecvMsg(reply)
if err == io.EOF {
return nil
}
if err == nil {
return status.Error(codes.Internal, "cardinality violation: expected <EOF> for non server-streaming RPCs, but received another message")
}
return err
}

View file

@ -553,6 +553,7 @@ func chainStreamClientInterceptors(cc *ClientConn) {
if cc.dopts.streamInt != nil {
interceptors = append([]StreamClientInterceptor{cc.dopts.streamInt}, interceptors...)
}
interceptors = append(interceptors, defaultStreamInterceptor)
var chainedInt StreamClientInterceptor
if len(interceptors) == 0 {
chainedInt = nil

View file

@ -27,32 +27,29 @@ import (
// extra goroutines. This is typically used for passing updates from one entity
// to another within gRPC.
//
// To avoid extra memory allocations and type assertions, using any on
// performance-critical code paths is discouraged. Use concrete types wherever
// possible when instantiating Unbounded.
//
// All methods on this type are thread-safe and don't block on anything except
// the underlying mutex used for synchronization.
//
// Unbounded supports values of any type to be stored in it by using a channel
// of `any`. This means that a call to Put() incurs an extra memory allocation,
// and also that users need a type assertion while reading. For performance
// critical code paths, using Unbounded is strongly discouraged and defining a
// new type specific implementation of this buffer is preferred. See
// internal/transport/transport.go for an example of this.
type Unbounded struct {
c chan any
type Unbounded[T any] struct {
c chan T
closed bool
closing bool
mu sync.Mutex
backlog []any
backlog []T
}
// NewUnbounded returns a new instance of Unbounded.
func NewUnbounded() *Unbounded {
return &Unbounded{c: make(chan any, 1)}
func NewUnbounded[T any]() *Unbounded[T] {
return &Unbounded[T]{c: make(chan T, 1)}
}
var errBufferClosed = errors.New("Put called on closed buffer.Unbounded")
// Put adds t to the unbounded buffer.
func (b *Unbounded) Put(t any) error {
func (b *Unbounded[T]) Put(t T) error {
b.mu.Lock()
defer b.mu.Unlock()
if b.closing {
@ -72,13 +69,14 @@ func (b *Unbounded) Put(t any) error {
// Load sends the earliest buffered data, if any, onto the read channel returned
// by Get(). Users are expected to call this every time they successfully read a
// value from the read channel.
func (b *Unbounded) Load() {
func (b *Unbounded[T]) Load() {
b.mu.Lock()
defer b.mu.Unlock()
if len(b.backlog) > 0 {
select {
case b.c <- b.backlog[0]:
b.backlog[0] = nil
var zero T
b.backlog[0] = zero
b.backlog = b.backlog[1:]
default:
}
@ -96,14 +94,14 @@ func (b *Unbounded) Load() {
//
// If the unbounded buffer is closed, the read channel returned by this method
// is closed after all data is drained.
func (b *Unbounded) Get() <-chan any {
func (b *Unbounded[T]) Get() <-chan T {
return b.c
}
// Close closes the unbounded buffer. No subsequent data may be Put(), and the
// channel returned from Get() will be closed after all the data is read and
// Load() is called for the final time.
func (b *Unbounded) Close() {
func (b *Unbounded[T]) Close() {
b.mu.Lock()
defer b.mu.Unlock()
if b.closing {

View file

@ -153,6 +153,13 @@ var (
// TODO: Remove this env var once v1.83.0 is released.
ControlBufferThrottleLimit = uint64FromEnv("GRPC_GO_EXPERIMENTAL_CONTROL_BUFFER_THROTTLE_LIMIT", 100, 1, 10000)
// ALTSMaxFrameSize is the maximum frame size for ALTS in bytes.
// This can be overridden by setting the environment variable
// "GRPC_GO_EXPERIMENTAL_ALTS_MAX_FRAME_SIZE" (min: 4096, max: 512*1024, capped at
// 512KiB to match altsWriteBufferMaxSize in
// credentials/alts/internal/conn/record.go).
ALTSMaxFrameSize = uint64FromEnv("GRPC_GO_EXPERIMENTAL_ALTS_MAX_FRAME_SIZE", 4096, 4096, 512*1024)
// EnableReceiveBufferCompaction enables the compaction of data buffers
// to reduce the number of buffers in the receive buffer.
//

View file

@ -98,4 +98,9 @@ var (
// filter. For more details, see:
// https://github.com/grpc/proposal/blob/master/A83-xds-gcp-authn-filter.md
GCPAuthenticationFilterEnabled = boolFromEnv("GRPC_EXPERIMENTAL_XDS_GCP_AUTHENTICATION_FILTER", false)
// XDSClientExtAuthzEnabled indicates whether the external authorization
// filter is enabled on the client side. For more details, see:
// https://github.com/grpc/proposal/blob/master/A92-xds-ext-authz.md
XDSClientExtAuthzEnabled = boolFromEnv("GRPC_EXPERIMENTAL_XDS_EXT_AUTHZ_ON_CLIENT", false)
)

View file

@ -41,7 +41,7 @@ type CallbackSerializer struct {
// its resources.
done chan struct{}
callbacks *buffer.Unbounded
callbacks *buffer.Unbounded[func(context.Context)]
}
// NewCallbackSerializer returns a new CallbackSerializer instance. The provided
@ -52,7 +52,7 @@ type CallbackSerializer struct {
func NewCallbackSerializer(ctx context.Context) *CallbackSerializer {
cs := &CallbackSerializer{
done: make(chan struct{}),
callbacks: buffer.NewUnbounded(),
callbacks: buffer.NewUnbounded[func(context.Context)](),
}
go cs.run(ctx)
return cs
@ -113,7 +113,7 @@ func (cs *CallbackSerializer) run(ctx context.Context) {
// Run all callbacks.
for cb := range cs.callbacks.Get() {
cs.callbacks.Load()
cb.(func(context.Context))(ctx)
cb(ctx)
}
}

View file

@ -0,0 +1,109 @@
/*
*
* Copyright 2026 gRPC authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*
*/
package grpcsync
import (
"sync/atomic"
"google.golang.org/grpc/grpclog"
)
var logger = grpclog.Component("grpcsync")
// RefCounted is a reference counted wrapper of type T. It tracks the number of
// active references and runs a cleanup when the last reference is released.
type RefCounted[T any] struct {
val T
refCount atomic.Int32
onZero func()
}
// NewRefCounted creates a new RefCounted instance wrapping the given value with
// initial refcount of one.
//
// The value should typically be a pointer, interface, or handle rather than a
// plain value type (such as a struct or primitive value).
//
// The provided onZero callback must not be nil, and is executed exactly once
// when the reference count drops to zero. Panics if onZero is nil.
//
// WARNING: onZero runs synchronously inside Decrement; it must not acquire
// locks held by Decrement callers.
func NewRefCounted[T any](val T, onZero func()) *RefCounted[T] {
if onZero == nil {
panic("grpcsync: onZero callback cannot be nil")
}
rc := &RefCounted[T]{
val: val,
onZero: onZero,
}
rc.refCount.Store(1)
return rc
}
// Value returns the encapsulated resource.
func (rc *RefCounted[T]) Value() T {
return rc.val
}
// TryIncrement attempts to increment the reference count, returning true if
// successful. It returns false if the count has already reached 0, indicating
// the resource has been cleaned up and cannot be resurrected.
//
// WARNING: Avoid calling TryIncrement on hot paths when an active reference is
// already guaranteed by the caller; use Increment instead to bypass the
// CompareAndSwap loop overhead. TryIncrement should be reserved for speculative
// lookups (such as fetching from a cache or map) where the resource might be
// dead.
func (rc *RefCounted[T]) TryIncrement() bool {
// Utilize a CompareAndSwap loop to prevent race conditions where a concurrent
// decrement could drop the count to zero between the read and the increment
// operation, which would otherwise inadvertently resurrect a closed resource.
for {
count := rc.refCount.Load()
if count <= 0 {
return false // Already dead or dying
}
if rc.refCount.CompareAndSwap(count, count+1) {
return true
}
}
}
// Increment increments the reference count.
//
// WARNING: Call Increment only when there is a guarantee that an active
// reference is already present, ensuring the resource is not dead. If there is
// a possibility that the resource is dead or its reference count might have
// reached zero, call TryIncrement instead.
func (rc *RefCounted[T]) Increment() {
if rc.refCount.Add(1) <= 1 {
logger.Errorf("Resource already closed or dead")
}
}
// Decrement decrements the reference count. If it drops to zero, the onZero
// callback is executed synchronously before this method returns.
func (rc *RefCounted[T]) Decrement() {
if v := rc.refCount.Add(-1); v < 0 {
logger.Errorf("Refcount cannot be negative")
} else if v == 0 {
rc.onZero()
}
}

View file

@ -245,6 +245,10 @@ var (
AsyncReporterCleanupDelegate = func(cleanup func()) func() {
return cleanup
}
// XDSFilterWrapperOption returns a ServerOption that sets the internal
// stream wrapper used for server-side xDS HTTP filters.
XDSFilterWrapperOption any // func(func(grpc.ServerStream) (grpc.ServerStream, error)) grpc.ServerOption
)
// HealthChecker defines the signature of the client-side LB channel health

View file

@ -42,6 +42,10 @@ type RPCInfo struct {
// efficiency reasons. SelectConfig should not be blocking.
Context context.Context
Method string // i.e. "/Service/Method"
// Authority is the target authority (host) name of the RPC. This is required
// for HTTP filters (such as external processing) to populate request
// attributes.
Authority string
}
// RPCConfig describes the configuration to use for each RPC.
@ -54,17 +58,6 @@ type RPCConfig struct {
Interceptor any
}
// ServerInterceptor is an interceptor for incoming RPC's on gRPC server side.
type ServerInterceptor interface {
// AllowRPC checks if an incoming RPC is allowed to proceed based on
// information about connection RPC was received on, and HTTP Headers. This
// information will be piped into context.
AllowRPC(ctx context.Context) error // TODO: Make this a real interceptor for filters such as rate limiting.
// Close closes the interceptor. Once called, no new calls to NewStream are
// accepted. Ongoing calls to NewStream are allowed to complete.
Close()
}
type csKeyType string
const csKey = csKeyType("grpc.internal.resolver.configSelector")

View file

@ -219,6 +219,17 @@ type outFlowControlSizeRequest struct {
resp chan uint32
}
// outStreamRequestForTesting is used by tests to retrieve a pointer to an outStream
// safely inside the loopyWriter goroutine without racing on loopyWriter.estdStreams.
// By holding the *outStream pointer while streams and transports close down, tests
// can inspect settled accounting values (such as bytesOutStanding) after loopyWriter
// has run to completion without any data races.
type outStreamRequestForTesting struct {
throttledItem
streamID uint32
resp chan *outStream
}
// closeConnection is an instruction to tell the loopy writer to flush the
// framer and exit, which will cause the transport's connection to be closed
// (by the client or server). The transport itself will close after the reader
@ -799,6 +810,10 @@ func (l *loopyWriter) outFlowControlSizeRequestHandler(o *outFlowControlSizeRequ
o.resp <- l.sendQuota
}
func (l *loopyWriter) outStreamRequestHandler(o *outStreamRequestForTesting) {
o.resp <- l.estdStreams[o.streamID]
}
func (l *loopyWriter) cleanupStreamHandler(c *cleanupStream) error {
c.onWrite()
if str, ok := l.estdStreams[c.streamID]; ok {
@ -896,6 +911,8 @@ func (l *loopyWriter) handle(i any) error {
return l.goAwayHandler(i)
case *outFlowControlSizeRequest:
l.outFlowControlSizeRequestHandler(i)
case *outStreamRequestForTesting:
l.outStreamRequestHandler(i)
case closeConnection:
// Just return a non-I/O error and run() will flush and close the
// connection.

View file

@ -554,8 +554,6 @@ func (t *http2Client) createHeaderFields(ctx context.Context, callHdr *CallHdr)
if err != nil {
return nil, err
}
// TODO(mmukhi): Benchmark if the performance gets better if count the metadata and other header fields
// first and create a slice of that exact size.
// Make the slice of certain predictable size to reduce allocations made by append.
hfLen := 7 // :method, :scheme, :path, :authority, content-type, user-agent, te
hfLen += len(authData) + len(callAuthData)
@ -575,6 +573,21 @@ func (t *http2Client) createHeaderFields(ctx context.Context, callHdr *CallHdr)
if _, ok := ctx.Deadline(); ok {
hfLen++
}
// Count the metadata header fields as well so the slice is not reallocated
// while they are appended below. Reserved headers are dropped when writing,
// so this may slightly over-count, which is preferable to growing the slice.
md, added, mdOK := metadataFromOutgoingContextRaw(ctx)
if mdOK {
for _, vv := range md {
hfLen += len(vv)
}
for _, vv := range added {
hfLen += len(vv) / 2
}
}
for _, vv := range t.md {
hfLen += len(vv)
}
headerFields := make([]hpack.HeaderField, 0, hfLen)
headerFields = append(headerFields, hpack.HeaderField{Name: ":method", Value: "POST"})
headerFields = append(headerFields, hpack.HeaderField{Name: ":scheme", Value: t.scheme})
@ -619,7 +632,7 @@ func (t *http2Client) createHeaderFields(ctx context.Context, callHdr *CallHdr)
headerFields = append(headerFields, hpack.HeaderField{Name: k, Value: encodeMetadataHeader(k, v)})
}
if md, added, ok := metadataFromOutgoingContextRaw(ctx); ok {
if mdOK {
var k string
for k, vv := range md {
// HTTP doesn't allow you to set pseudoheaders after non pseudoheaders were set.
@ -648,6 +661,9 @@ func (t *http2Client) createHeaderFields(ctx context.Context, callHdr *CallHdr)
if isReservedHeader(k) {
continue
}
if err := imetadata.ValidatePair(k, vv...); err != nil {
return nil, status.Error(codes.Internal, err.Error())
}
for _, v := range vv {
headerFields = append(headerFields, hpack.HeaderField{Name: k, Value: encodeMetadataHeader(k, v)})
}
@ -691,6 +707,9 @@ func (t *http2Client) getTrAuthData(ctx context.Context, audience string) (map[s
for k, v := range data {
// Capital header names are illegal in HTTP/2.
k = strings.ToLower(k)
if err := imetadata.ValidatePair(k, v); err != nil {
return nil, status.Error(codes.Internal, err.Error())
}
authData[k] = v
}
}
@ -724,6 +743,9 @@ func (t *http2Client) getCallAuthData(ctx context.Context, audience string, call
for k, v := range data {
// Capital header names are illegal in HTTP/2
k = strings.ToLower(k)
if err := imetadata.ValidatePair(k, v); err != nil {
return nil, status.Error(codes.Internal, err.Error())
}
callAuthData[k] = v
}
}
@ -1226,23 +1248,26 @@ func (t *http2Client) handleData(f *parsedDataFrame) {
t.closeStream(s, io.EOF, true, http2.ErrCodeFlowControl, status.New(codes.Internal, err.Error()), nil, false)
return
}
}
if s.nonGRPCStatus != nil {
// The frame should be handled as a non-gRPC response body
st := s.handleNonGRPCData(f)
if st != nil {
t.closeStream(s, st.Err(), true, http2.ErrCodeProtocol, st, nil, true)
return
}
if w := s.fc.onRead(size); w > 0 {
t.controlBuf.put(&outgoingWindowUpdate{
streamID: s.id,
increment: w,
})
}
if s.nonGRPCStatus != nil {
// The frame should be handled as a non-gRPC response body. A non-nil
// status also covers END_STREAM, so no separate handling is needed below.
st := s.handleNonGRPCData(f)
if st != nil {
t.closeStream(s, st.Err(), true, http2.ErrCodeProtocol, st, nil, true)
return
}
if w := s.fc.onRead(size); w > 0 {
t.controlBuf.put(&outgoingWindowUpdate{
streamID: s.id,
increment: w,
})
}
return
}
if size > 0 {
dataLen := f.data.Len()
if f.Header().Flags.Has(http2.FlagDataPadded) {
if w := s.fc.onRead(size - uint32(dataLen)); w > 0 {

View file

@ -531,6 +531,7 @@ func (t *http2Server) operateHeaders(ctx context.Context, frame *http2.MetaHeade
if frame.StreamEnded() {
// s is just created by the caller. No lock needed.
s.state = streamReadDone
s.write(recvMsg{err: io.EOF})
}
if timeoutSet {
s.ctx, s.cancel = context.WithTimeout(ctx, timeout)

View file

@ -300,24 +300,29 @@ func decodeGrpcMessageUnchecked(msg string) string {
type bufWriter struct {
pool *imem.SimpleBufferPool
buf []byte
bufHandle *[]byte
offset int
batchSize int
conn io.Writer
err error
}
// unsharedBufWriter keeps the unshared slice header in the writer allocation.
type unsharedBufWriter struct {
bufWriter
buf []byte
}
func newBufWriter(conn io.Writer, batchSize int, pool *imem.SimpleBufferPool) *bufWriter {
w := &bufWriter{
batchSize: batchSize,
conn: conn,
pool: pool,
if pool == nil && batchSize > 0 {
w := &unsharedBufWriter{
bufWriter: bufWriter{batchSize: batchSize, conn: conn},
buf: make([]byte, batchSize),
}
w.bufHandle = &w.buf
return &w.bufWriter
}
// this indicates that we should use non shared buf
if pool == nil {
w.buf = make([]byte, batchSize)
}
return w
return &bufWriter{batchSize: batchSize, conn: conn, pool: pool}
}
func (w *bufWriter) Write(b []byte) (int, error) {
@ -328,13 +333,13 @@ func (w *bufWriter) Write(b []byte) (int, error) {
n, err := w.conn.Write(b)
return n, toIOError(err)
}
if w.buf == nil {
b := w.pool.Get(w.batchSize)
w.buf = *b
if w.bufHandle == nil {
w.bufHandle = w.pool.Get(w.batchSize)
}
buf := *w.bufHandle
written := 0
for len(b) > 0 {
copied := copy(w.buf[w.offset:], b)
copied := copy(buf[w.offset:], b)
b = b[copied:]
written += copied
w.offset += copied
@ -350,15 +355,17 @@ func (w *bufWriter) Write(b []byte) (int, error) {
func (w *bufWriter) Flush() error {
err := w.flushKeepBuffer()
// Only release the buffer if we are in a "shared" mode
if w.buf != nil && w.pool != nil {
b := w.buf
w.pool.Put(&b)
w.buf = nil
}
w.releaseBuffer()
return err
}
func (w *bufWriter) releaseBuffer() {
if w.pool != nil && w.bufHandle != nil {
w.pool.Put(w.bufHandle)
w.bufHandle = nil
}
}
func (w *bufWriter) flushKeepBuffer() error {
if w.err != nil {
return w.err
@ -366,9 +373,13 @@ func (w *bufWriter) flushKeepBuffer() error {
if w.offset == 0 {
return nil
}
_, w.err = w.conn.Write(w.buf[:w.offset])
buf := *w.bufHandle
_, w.err = w.conn.Write(buf[:w.offset])
w.err = toIOError(w.err)
w.offset = 0
if w.err != nil {
w.releaseBuffer()
}
return w.err
}

View file

@ -89,10 +89,10 @@ type recvMsg struct {
// recvBuffer is an unbounded channel of recvMsg structs.
//
// Note: recvBuffer differs from buffer.Unbounded only in the fact that it
// holds a channel of recvMsg structs instead of objects implementing "item"
// interface. recvBuffer is written to much more often and using strict recvMsg
// structs helps avoid allocation in "recvBuffer.put"
// Note: recvBuffer differs from buffer.Unbounded in that it provides in-place
// value initialization (via init) to avoid struct pointer allocations on
// streams, and automatically frees pooled mem.Buffer payloads when stream
// errors occur.
type recvBuffer struct {
c chan recvMsg
mu sync.Mutex
@ -466,15 +466,18 @@ func (s *Stream) ReadMessageHeader(header []byte) (err error) {
return er
}
s.readRequester.requestRead(len(header))
bytesRead := 0
for len(header) != 0 {
n, err := s.trReader.ReadMessageHeader(header)
bytesRead += n
header = header[n:]
if len(header) == 0 {
err = nil
}
if err != nil {
if n > 0 && err == io.EOF {
if bytesRead > 0 && err == io.EOF {
err = io.ErrUnexpectedEOF
s.trReader.er = err
}
return err
}
@ -505,19 +508,22 @@ func (s *Stream) read(n int) (data mem.BufferSlice, err error) {
allocCap := min(ceil(n, http2MaxFrameLen), 128)
data = make(mem.BufferSlice, 0, allocCap)
s.readRequester.requestRead(n)
bytesRead := 0
for n != 0 {
buf, err := s.trReader.Read(n)
var bufLen int
if buf != nil {
bufLen = buf.Len()
}
bytesRead += bufLen
n -= bufLen
if n == 0 {
err = nil
}
if err != nil {
if bufLen > 0 && err == io.EOF {
if bytesRead > 0 && err == io.EOF {
err = io.ErrUnexpectedEOF
s.trReader.er = err
}
data.Free()
return nil, err

View file

@ -24,6 +24,7 @@ package metadata // import "google.golang.org/grpc/metadata"
import (
"context"
"fmt"
"sort"
"strings"
"google.golang.org/grpc/internal"
@ -90,6 +91,76 @@ func Pairs(kv ...string) MD {
return md
}
// loggableMetadataKeys is the set of metadata keys whose values are known not
// to carry credentials or other sensitive data, and are therefore safe to print
// verbatim in String. Values for any key not in this set are redacted. Keys are
// in metadata's canonical lowercase form.
//
// The list is intentionally restricted to standardized gRPC and HTTP/2 protocol
// headers. Everything else, including all application-defined keys, is
// censored.
// Note that this list might change as we add/remove support for
// metadata types.
var loggableMetadataKeys = map[string]bool{
"content-type": true,
"te": true,
"user-agent": true,
"grpc-encoding": true,
"grpc-accept-encoding": true,
"grpc-timeout": true,
"grpc-status": true,
"grpc-message-type": true,
"grpc-previous-rpc-attempts": true,
"grpc-retry-pushback-ms": true,
}
// String implements fmt.Stringer to allow metadata to be printed when stored in
// a context.
//
// To avoid accidentally leaking credentials or other sensitive data (for
// example via log(md), or by logging a value that happens to contain metadata),
// String reports only keys on an allowlist of standardized, non-sensitive
// protocol headers (see loggableMetadataKeys). Every other key is omitted
// entirely, along with its values, because a key name may itself be sensitive;
// only the number of omitted keys is reported, as "<N redacted>". This is a
// best-effort guard against accidents, not a security boundary: a caller that
// genuinely wants the full contents can still print map[string][]string(md)
// directly.
//
// Note that this only affects verbs that use the Stringer, such as %v and %s.
// The %#v verb prints the underlying map with all values and is not redacted.
// Users should not rely on the output of this method to be stable
// # Experimental
//
// Notice: This API is EXPERIMENTAL and may be changed or removed in a later
// release.
func (md MD) String() string {
keys := make([]string, 0, len(md))
for k := range md {
if loggableMetadataKeys[strings.ToLower(k)] {
keys = append(keys, k)
}
}
sort.Strings(keys)
var sb strings.Builder
sb.WriteString("map[")
for i, k := range keys {
if i > 0 {
sb.WriteByte(' ')
}
fmt.Fprintf(&sb, "%s:%v", k, md[k])
}
if redacted := len(md) - len(keys); redacted > 0 {
if len(keys) > 0 {
sb.WriteByte(' ')
}
fmt.Fprintf(&sb, "<%d redacted>", redacted)
}
sb.WriteByte(']')
return sb.String()
}
// Len returns the number of items in md.
func (md MD) Len() int {
return len(md)

View file

@ -188,7 +188,7 @@ func (ccr *ccResolverWrapper) ParseServiceConfig(scJSON string) *serviceconfig.P
// addChannelzTraceEvent adds a channelz trace event containing the new
// state received from resolver implementations.
func (ccr *ccResolverWrapper) addChannelzTraceEvent(s resolver.State) {
if !logger.V(0) && !channelz.IsOn() {
if !logger.V(2) && !channelz.IsOn() {
return
}
var updates []string

View file

@ -89,6 +89,7 @@ func init() {
internal.MetricsRecorderForServer = func(srv *Server) estats.MetricsRecorder {
return istats.NewMetricsRecorderList(srv.opts.statsHandlers)
}
internal.XDSFilterWrapperOption = xdsFilterWrapperOption
}
var statusOK = status.New(codes.OK, "")
@ -117,10 +118,8 @@ type ServiceDesc struct {
// serviceInfo wraps information about a service. It is very similar to
// ServiceDesc and is constructed from it for internal purposes.
type serviceInfo struct {
// Contains the implementation for the methods in this service.
serviceImpl any
methods map[string]*MethodDesc
streams map[string]*StreamDesc
serviceImpl any // Implementation for the methods in this service.
streams map[string]*StreamDesc // Streaming descriptors and wrapped unary descriptors for this service.
mdata any
}
@ -183,6 +182,7 @@ type serverOptions struct {
bufferPool mem.BufferPool
waitForHandlers bool
staticWindowSize bool
streamWrapper func(ServerStream) (ServerStream, error)
}
var defaultServerOptions = serverOptions{
@ -663,6 +663,14 @@ func bufferPool(bufferPool mem.BufferPool) ServerOption {
})
}
// xdsFilterWrapperOption returns a ServerOption that sets the server-level
// stream wrapper (used internally by xDS server filters).
func xdsFilterWrapperOption(w func(ServerStream) (ServerStream, error)) ServerOption {
return newFuncServerOption(func(o *serverOptions) {
o.streamWrapper = w
})
}
// serverWorkerResetThreshold defines how often the stack must be reset. Every
// N requests, by spawning a new goroutine in its place, a worker can reset its
// stack so that large stacks don't live in memory forever. 2^16 should allow
@ -790,18 +798,22 @@ func (s *Server) register(sd *ServiceDesc, ss any) {
}
info := &serviceInfo{
serviceImpl: ss,
methods: make(map[string]*MethodDesc),
streams: make(map[string]*StreamDesc),
mdata: sd.Metadata,
}
for i := range sd.Methods {
d := &sd.Methods[i]
info.methods[d.MethodName] = d
}
for i := range sd.Streams {
d := &sd.Streams[i]
info.streams[d.StreamName] = d
}
for i := range sd.Methods {
d := &sd.Methods[i]
info.streams[d.MethodName] = &StreamDesc{
StreamName: d.MethodName,
Handler: s.wrapUnaryHandler(d),
ServerStreams: false,
ClientStreams: false,
}
}
s.services[sd.ServiceName] = info
}
@ -827,20 +839,26 @@ type ServiceInfo struct {
func (s *Server) GetServiceInfo() map[string]ServiceInfo {
ret := make(map[string]ServiceInfo)
for n, srv := range s.services {
methods := make([]MethodInfo, 0, len(srv.methods)+len(srv.streams))
for m := range srv.methods {
methods = append(methods, MethodInfo{
Name: m,
IsClientStream: false,
IsServerStream: false,
})
}
methods := make([]MethodInfo, 0, len(srv.streams))
// Iterate over unary methods first to maintain backward compatibility of order.
for m, d := range srv.streams {
methods = append(methods, MethodInfo{
Name: m,
IsClientStream: d.ClientStreams,
IsServerStream: d.ServerStreams,
})
if !d.ClientStreams && !d.ServerStreams {
methods = append(methods, MethodInfo{
Name: m,
IsClientStream: false,
IsServerStream: false,
})
}
}
// Iterate over streaming methods.
for m, d := range srv.streams {
if d.ClientStreams || d.ServerStreams {
methods = append(methods, MethodInfo{
Name: m,
IsClientStream: d.ClientStreams,
IsServerStream: d.ServerStreams,
})
}
}
ret[n] = ServiceInfo{
@ -1185,42 +1203,6 @@ func (s *Server) incrCallsFailed() {
s.channelz.ServerMetrics.CallsFailed.Add(1)
}
func (s *Server) sendResponse(ctx context.Context, stream *transport.ServerStream, msg any, cp Compressor, opts *transport.WriteOptions, comp encoding.Compressor) error {
data, err := encode(s.getCodec(stream.ContentSubtype()), msg)
if err != nil {
channelz.Error(logger, s.channelz, "grpc: server failed to encode response: ", err)
return err
}
compData, pf, err := compress(data, cp, comp, s.opts.bufferPool)
if err != nil {
data.Free()
channelz.Error(logger, s.channelz, "grpc: server failed to compress response: ", err)
return err
}
hdr, payload := msgHeader(data, compData, pf)
defer func() {
compData.Free()
data.Free()
// payload does not need to be freed here, it is either data or compData, both of
// which are already freed.
}()
dataLen := data.Len()
payloadLen := payload.Len()
// TODO(dfawley): should we be checking len(data) instead?
if payloadLen > s.opts.maxSendMessageSize {
return status.Errorf(codes.ResourceExhausted, "grpc: trying to send message larger than max (%d vs. %d)", payloadLen, s.opts.maxSendMessageSize)
}
err = stream.Write(hdr, payload, opts)
if err == nil && s.statsHandler != nil {
s.statsHandler.HandleRPC(ctx, outPayload(false, msg, dataLen, payloadLen, time.Now()))
}
return err
}
// chainUnaryServerInterceptors chains all unary server interceptors into one.
func chainUnaryServerInterceptors(s *Server) {
// Prepend opts.unaryInt to the chaining interceptors if it exists, since unaryInt will
@ -1257,300 +1239,6 @@ func getChainUnaryHandler(interceptors []UnaryServerInterceptor, curr int, info
}
}
func (s *Server) processUnaryRPC(ctx context.Context, stream *transport.ServerStream, info *serviceInfo, md *MethodDesc, trInfo *traceInfo) (err error) {
sh := s.statsHandler
if sh != nil || trInfo != nil || channelz.IsOn() {
if channelz.IsOn() {
s.incrCallsStarted()
}
var statsBegin *stats.Begin
if sh != nil {
statsBegin = &stats.Begin{
BeginTime: time.Now(),
IsClientStream: false,
IsServerStream: false,
}
sh.HandleRPC(ctx, statsBegin)
}
if trInfo != nil {
trInfo.tr.LazyLog(&trInfo.firstLine, false)
}
// The deferred error handling for tracing, stats handler and channelz are
// combined into one function to reduce stack usage -- a defer takes ~56-64
// bytes on the stack, so overflowing the stack will require a stack
// re-allocation, which is expensive.
//
// To maintain behavior similar to separate deferred statements, statements
// should be executed in the reverse order. That is, tracing first, stats
// handler second, and channelz last. Note that panics *within* defers will
// lead to different behavior, but that's an acceptable compromise; that
// would be undefined behavior territory anyway.
defer func() {
if trInfo != nil {
if err != nil && err != io.EOF {
trInfo.tr.LazyLog(&fmtStringer{"%v", []any{err}}, true)
trInfo.tr.SetError()
}
trInfo.tr.Finish()
}
if sh != nil {
end := &stats.End{
BeginTime: statsBegin.BeginTime,
EndTime: time.Now(),
}
if err != nil && err != io.EOF {
end.Error = toRPCErr(err)
}
sh.HandleRPC(ctx, end)
}
if channelz.IsOn() {
if err != nil && err != io.EOF {
s.incrCallsFailed()
} else {
s.incrCallsSucceeded()
}
}
}()
}
var binlogs []binarylog.MethodLogger
if ml := binarylog.GetMethodLogger(stream.Method()); ml != nil {
binlogs = append(binlogs, ml)
}
if s.opts.binaryLogger != nil {
if ml := s.opts.binaryLogger.GetMethodLogger(stream.Method()); ml != nil {
binlogs = append(binlogs, ml)
}
}
if len(binlogs) != 0 {
md, _ := metadata.FromIncomingContext(ctx)
logEntry := &binarylog.ClientHeader{
Header: md,
MethodName: stream.Method(),
PeerAddr: nil,
}
if deadline, ok := ctx.Deadline(); ok {
logEntry.Timeout = time.Until(deadline)
if logEntry.Timeout < 0 {
logEntry.Timeout = 0
}
}
if a := md[":authority"]; len(a) > 0 {
logEntry.Authority = a[0]
}
if peer, ok := peer.FromContext(ctx); ok {
logEntry.PeerAddr = peer.Addr
}
for _, binlog := range binlogs {
binlog.Log(ctx, logEntry)
}
}
// comp and cp are used for compression. decomp and dc are used for
// decompression. If comp and decomp are both set, they are the same;
// however they are kept separate to ensure that at most one of the
// compressor/decompressor variable pairs are set for use later.
var comp, decomp encoding.Compressor
var cp Compressor
var dc Decompressor
var sendCompressorName string
// If dc is set and matches the stream's compression, use it. Otherwise, try
// to find a matching registered compressor for decomp.
if rc := stream.RecvCompress(); s.opts.dc != nil && s.opts.dc.Type() == rc {
dc = s.opts.dc
} else if rc != "" && rc != encoding.Identity {
decomp = encoding.GetCompressor(rc)
if decomp == nil {
st := status.Newf(codes.Unimplemented, "grpc: Decompressor is not installed for grpc-encoding %q", rc)
stream.WriteStatus(st)
return st.Err()
}
}
// If cp is set, use it. Otherwise, attempt to compress the response using
// the incoming message compression method.
//
// NOTE: this needs to be ahead of all handling, https://github.com/grpc/grpc-go/issues/686.
if s.opts.cp != nil {
cp = s.opts.cp
sendCompressorName = cp.Type()
} else if rc := stream.RecvCompress(); rc != "" && rc != encoding.Identity {
// Legacy compressor not specified; attempt to respond with same encoding.
comp = encoding.GetCompressor(rc)
if comp != nil {
sendCompressorName = comp.Name()
}
}
if sendCompressorName != "" {
if err := stream.SetSendCompress(sendCompressorName); err != nil {
return status.Errorf(codes.Internal, "grpc: failed to set send compressor: %v", err)
}
}
var payInfo *payloadInfo
if sh != nil || len(binlogs) != 0 {
payInfo = &payloadInfo{}
defer payInfo.free()
}
d, err := recvAndDecompress(&parser{r: stream, bufferPool: s.opts.bufferPool}, stream, dc, s.opts.maxReceiveMessageSize, payInfo, decomp, true)
if err != nil {
if e := stream.WriteStatus(status.Convert(err)); e != nil {
channelz.Warningf(logger, s.channelz, "grpc: Server.processUnaryRPC failed to write status: %v", e)
}
return err
}
freed := false
dataFree := func() {
if !freed {
d.Free()
freed = true
}
}
defer dataFree()
df := func(v any) error {
defer dataFree()
if err := s.getCodec(stream.ContentSubtype()).Unmarshal(d, v); err != nil {
return status.Errorf(codes.Internal, "grpc: error unmarshalling request: %v", err)
}
if sh != nil {
sh.HandleRPC(ctx, &stats.InPayload{
RecvTime: time.Now(),
Payload: v,
Length: d.Len(),
WireLength: payInfo.compressedLength + headerLen,
CompressedLength: payInfo.compressedLength,
})
}
if len(binlogs) != 0 {
cm := &binarylog.ClientMessage{
Message: d.Materialize(),
}
for _, binlog := range binlogs {
binlog.Log(ctx, cm)
}
}
if trInfo != nil {
trInfo.tr.LazyLog(&payload{sent: false, msg: v}, true)
}
return nil
}
ctx = NewContextWithServerTransportStream(ctx, stream)
reply, appErr := md.Handler(info.serviceImpl, ctx, df, s.opts.unaryInt)
if appErr != nil {
appStatus, ok := status.FromError(appErr)
if !ok {
// Convert non-status application error to a status error with code
// Unknown, but handle context errors specifically.
appStatus = status.FromContextError(appErr)
appErr = appStatus.Err()
}
if trInfo != nil {
trInfo.tr.LazyLog(stringer(appStatus.Message()), true)
trInfo.tr.SetError()
}
if e := stream.WriteStatus(appStatus); e != nil {
channelz.Warningf(logger, s.channelz, "grpc: Server.processUnaryRPC failed to write status: %v", e)
}
if len(binlogs) != 0 {
if h, _ := stream.Header(); h.Len() > 0 {
// Only log serverHeader if there was header. Otherwise it can
// be trailer only.
sh := &binarylog.ServerHeader{
Header: h,
}
for _, binlog := range binlogs {
binlog.Log(ctx, sh)
}
}
st := &binarylog.ServerTrailer{
Trailer: stream.Trailer(),
Err: appErr,
}
for _, binlog := range binlogs {
binlog.Log(ctx, st)
}
}
return appErr
}
if trInfo != nil {
trInfo.tr.LazyLog(stringer("OK"), false)
}
opts := &transport.WriteOptions{Last: true}
// Server handler could have set new compressor by calling SetSendCompressor.
// In case it is set, we need to use it for compressing outbound message.
if stream.SendCompress() != sendCompressorName {
comp = encoding.GetCompressor(stream.SendCompress())
}
if err := s.sendResponse(ctx, stream, reply, cp, opts, comp); err != nil {
if err == io.EOF {
// The entire stream is done (for unary RPC only).
return err
}
if sts, ok := status.FromError(err); ok {
if e := stream.WriteStatus(sts); e != nil {
channelz.Warningf(logger, s.channelz, "grpc: Server.processUnaryRPC failed to write status: %v", e)
}
} else {
switch st := err.(type) {
case transport.ConnectionError:
// Nothing to do here.
default:
panic(fmt.Sprintf("grpc: Unexpected error (%T) from sendResponse: %v", st, st))
}
}
if len(binlogs) != 0 {
h, _ := stream.Header()
sh := &binarylog.ServerHeader{
Header: h,
}
st := &binarylog.ServerTrailer{
Trailer: stream.Trailer(),
Err: appErr,
}
for _, binlog := range binlogs {
binlog.Log(ctx, sh)
binlog.Log(ctx, st)
}
}
return err
}
if len(binlogs) != 0 {
h, _ := stream.Header()
sh := &binarylog.ServerHeader{
Header: h,
}
sm := &binarylog.ServerMessage{
Message: reply,
}
for _, binlog := range binlogs {
binlog.Log(ctx, sh)
binlog.Log(ctx, sm)
}
}
if trInfo != nil {
trInfo.tr.LazyLog(&payload{sent: true, msg: reply}, true)
}
// TODO: Should we be logging if writing status failed here, like above?
// Should the logging be in WriteStatus? Should we ignore the WriteStatus
// error or allow the stats handler to see it?
if len(binlogs) != 0 {
st := &binarylog.ServerTrailer{
Trailer: stream.Trailer(),
Err: appErr,
}
for _, binlog := range binlogs {
binlog.Log(ctx, st)
}
}
return stream.WriteStatus(statusOK)
}
// chainStreamServerInterceptors chains all stream server interceptors into one.
func chainStreamServerInterceptors(s *Server) {
// Prepend opts.streamInt to the chaining interceptors if it exists, since streamInt will
@ -1587,7 +1275,21 @@ func getChainStreamHandler(interceptors []StreamServerInterceptor, curr int, inf
}
}
func (s *Server) processStreamingRPC(ctx context.Context, stream *transport.ServerStream, info *serviceInfo, sd *StreamDesc, trInfo *traceInfo) (err error) {
// wrapUnaryHandler converts a standard Unary Method Descriptor into a
// StreamHandler, allowing Unary RPCs to be processed via the same unified
// pipeline as Streaming RPCs.
func (s *Server) wrapUnaryHandler(md *MethodDesc) StreamHandler {
return func(srv any, stream ServerStream) error {
df := func(v any) error { return stream.RecvMsg(v) }
reply, err := md.Handler(srv, stream.Context(), df, s.opts.unaryInt)
if err != nil {
return err
}
return stream.SendMsg(reply)
}
}
func (s *Server) processRPC(ctx context.Context, stream *transport.ServerStream, info *serviceInfo, sd *StreamDesc, trInfo *traceInfo) (err error) {
if channelz.IsOn() {
s.incrCallsStarted()
}
@ -1615,7 +1317,16 @@ func (s *Server) processStreamingRPC(ctx context.Context, stream *transport.Serv
}
if sh != nil || trInfo != nil || channelz.IsOn() {
// See comment in processUnaryRPC on defers.
// The deferred error handling for tracing, stats handler and channelz are
// combined into one function to reduce stack usage -- a defer takes ~56-64
// bytes on the stack, so overflowing the stack will require a stack
// re-allocation, which is expensive.
//
// To maintain behavior similar to separate deferred statements, statements
// should be executed in the reverse order. That is, tracing first, stats
// handler second, and channelz last. Note that panics *within* defers will
// lead to different behavior, but that's an acceptable compromise; that
// would be undefined behavior territory anyway.
defer func() {
if trInfo != nil {
ss.mu.Lock()
@ -1689,7 +1400,7 @@ func (s *Server) processStreamingRPC(ctx context.Context, stream *transport.Serv
ss.decompressorV1 = encoding.GetCompressor(rc)
if ss.decompressorV1 == nil {
st := status.Newf(codes.Unimplemented, "grpc: Decompressor is not installed for grpc-encoding %q", rc)
ss.s.WriteStatus(st)
ss.writeStatus(st)
return st.Err()
}
}
@ -1717,23 +1428,34 @@ func (s *Server) processStreamingRPC(ctx context.Context, stream *transport.Serv
ss.ctx = newContextWithRPCInfo(ss.ctx, false, ss.codec, ss.compressorV0, ss.compressorV1)
if trInfo != nil {
trInfo.tr.LazyLog(&trInfo.firstLine, false)
}
// Execute the xDS HTTP filter interceptors.
var appErr error
var server any
if info != nil {
server = info.serviceImpl
var wrappedStream ServerStream = ss
if s.opts.streamWrapper != nil {
wrappedStream, appErr = s.opts.streamWrapper(ss)
}
if s.opts.streamInt == nil {
appErr = sd.Handler(server, ss)
} else {
info := &StreamServerInfo{
FullMethod: stream.Method(),
IsClientStream: sd.ClientStreams,
IsServerStream: sd.ServerStreams,
if appErr == nil {
if trInfo != nil {
trInfo.tr.LazyLog(&trInfo.firstLine, false)
}
var server any
if info != nil {
server = info.serviceImpl
}
if s.opts.streamInt == nil || (!sd.ClientStreams && !sd.ServerStreams) {
// If there is no stream interceptor, or if this is a unary RPC, call
// the handler directly. The wrapped unary handler will call the unary
// interceptor if it exists.
appErr = sd.Handler(server, wrappedStream)
} else {
info := &StreamServerInfo{
FullMethod: stream.Method(),
IsClientStream: sd.ClientStreams,
IsServerStream: sd.ServerStreams,
}
appErr = s.opts.streamInt(server, wrappedStream, info, sd.Handler)
}
appErr = s.opts.streamInt(server, ss, info, sd.Handler)
}
if appErr != nil {
appStatus, ok := status.FromError(appErr)
@ -1750,6 +1472,17 @@ func (s *Server) processStreamingRPC(ctx context.Context, stream *transport.Serv
ss.mu.Unlock()
}
if len(ss.binlogs) != 0 {
if !ss.serverHeaderBinlogged {
if h, _ := ss.s.Header(); h.Len() > 0 {
sh := &binarylog.ServerHeader{
Header: h,
}
ss.serverHeaderBinlogged = true
for _, binlog := range ss.binlogs {
binlog.Log(ctx, sh)
}
}
}
st := &binarylog.ServerTrailer{
Trailer: ss.s.Trailer(),
Err: appErr,
@ -1758,8 +1491,9 @@ func (s *Server) processStreamingRPC(ctx context.Context, stream *transport.Serv
binlog.Log(ctx, st)
}
}
ss.s.WriteStatus(appStatus)
// TODO: Should we log an error from WriteStatus here and below?
if err := ss.writeStatus(appStatus); err != nil {
channelz.Warningf(logger, s.channelz, "grpc: Server.processRPC failed to write status: %v", err)
}
return appErr
}
if trInfo != nil {
@ -1776,7 +1510,10 @@ func (s *Server) processStreamingRPC(ctx context.Context, stream *transport.Serv
binlog.Log(ctx, st)
}
}
return ss.s.WriteStatus(statusOK)
if err := ss.writeStatus(statusOK); err != nil {
channelz.Warningf(logger, s.channelz, "grpc: Server.processRPC failed to write status: %v", err)
}
return err
}
func (s *Server) handleMalformedMethodName(stream *transport.ServerStream, ti *traceInfo) {
@ -1854,18 +1591,14 @@ func (s *Server) handleStream(t transport.ServerTransport, stream *transport.Ser
srv, knownService := s.services[service]
if knownService {
if md, ok := srv.methods[method]; ok {
s.processUnaryRPC(ctx, stream, srv, md, ti)
return
}
if sd, ok := srv.streams[method]; ok {
s.processStreamingRPC(ctx, stream, srv, sd, ti)
s.processRPC(ctx, stream, srv, sd, ti)
return
}
}
// Unknown service, or known server unknown method.
if unknownDesc := s.opts.unknownStreamDesc; unknownDesc != nil {
s.processStreamingRPC(ctx, stream, nil, unknownDesc, ti)
s.processRPC(ctx, stream, nil, unknownDesc, ti)
return
}
var errDesc string
@ -2069,9 +1802,6 @@ func (s *Server) isRegisteredMethod(serviceMethod string) bool {
method := serviceMethod[pos+1:]
srv, knownService := s.services[service]
if knownService {
if _, ok := srv.methods[method]; ok {
return true
}
if _, ok := srv.streams[method]; ok {
return true
}

View file

@ -28,6 +28,7 @@ import (
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
"google.golang.org/grpc/balancer"
@ -148,6 +149,94 @@ type ClientStream interface {
RecvMsg(m any) error
}
// clientStreamWrapper wraps a ClientStream and handles SendMsg, CloseSend, and
// RecvMsg parities based on the nature of stream.
type clientStreamWrapper struct {
ClientStream
desc *StreamDesc
closeSendCalled atomic.Bool
}
// CloseSend closes the send direction of the stream. The implementation ensures
// that CloseSend is only called once on the underlying ClientStream, even if
// CloseSend is called multiple times on the wrapper.
func (w *clientStreamWrapper) CloseSend() error {
if w.closeSendCalled.Swap(true) {
return nil
}
return w.ClientStream.CloseSend()
}
// SendMsg sends message m across the stream. For RPCs where client can call
// SendMsg only once, i.e. only server-streaming RPCs, it converts io.EOF to nil
// and immediately calls CloseSend to trigger any interceptor hooks.
func (w *clientStreamWrapper) SendMsg(m any) error {
err := w.ClientStream.SendMsg(m)
// If the RPC is a client-streaming RPC, the client can send multiple
// messages. In this case, the client should handle any type of
// error,including io.EOF and call CloseSend once it is done sending messages.
if w.desc.ClientStreams {
return err
}
if err == io.EOF {
// For non-client-streaming RPCs, we return nil instead of EOF on error
// because the generated code requires it. finish is not called; RecvMsg()
// will call it with the stream's status independently.
return nil
}
if err != nil {
return err
}
// In some scenarios (e.g., xDS), the same interceptors process both unary and
// streaming RPCs, relying on CloseSend to signal that no more messages are on
// the way. Although protobuf-generated stubs already invoke CloseSend for
// server-streaming RPCs, it is explicitly called here to ensure downstream
// interceptors are also notified when callers interact with the ClientStream
// API directly.
if err := w.CloseSend(); err != nil && err != io.EOF {
return err
}
return nil
}
// RecvMsg receives message m from the stream. For RPCs that call RecvMsg only
// once i.e. only client streaming RPCs, it calls the underlying RecvMsg a
// second time after receiving the first message to get the trailers.
func (w *clientStreamWrapper) RecvMsg(m any) error {
err := w.ClientStream.RecvMsg(m)
if err != nil {
return err
}
if w.desc.ServerStreams {
return nil
}
// Call RecvMsg again for non-server streaming RPCs to get the trailers and
// ensure RPC has completed successfully.
err = w.ClientStream.RecvMsg(m)
if err == io.EOF {
return nil
}
if err == nil {
return status.Error(codes.Internal, "cardinality violation: expected <EOF> for non server-streaming RPCs, but received another message")
}
return err
}
// defaultStreamInterceptor is a StreamClientInterceptor which wraps the
// ClientStream and is always invoked as the first interceptor. It consolidates
// behavior for different RPC types at the level closest to the application,
// which simplifies the underlying stream implementation and other interceptors
// by avoiding duplicate or scattered handling.
func defaultStreamInterceptor(ctx context.Context, desc *StreamDesc, cc *ClientConn, method string, streamer Streamer, opts ...CallOption) (ClientStream, error) {
cs, err := streamer(ctx, desc, cc, method, opts...)
if err != nil {
return nil, err
}
return &clientStreamWrapper{ClientStream: cs, desc: desc}, nil
}
// NewStream creates a new Stream for the client side. This is typically
// called by generated code. ctx is used for the lifetime of the stream.
//
@ -253,14 +342,11 @@ func newClientStream(ctx context.Context, desc *StreamDesc, cc *ClientConn, meth
mc := &emptyMethodConfig
var onCommit func()
newStream := func(ctx context.Context, filterOpts ...CallOption) (ClientStream, error) {
if filterOpts != nil {
opts = combine(opts, filterOpts)
}
newStream := func(ctx context.Context, opts ...CallOption) (ClientStream, error) {
return newClientStreamWithParams(ctx, desc, cc, method, mc, onCommit, nameResolutionDelayed, opts...)
}
rpcInfo := iresolver.RPCInfo{Context: ctx, Method: method}
rpcInfo := iresolver.RPCInfo{Context: ctx, Method: method, Authority: cc.authority}
rpcConfig, err := cc.safeConfigSelector.SelectConfig(rpcInfo)
if err != nil {
if st, ok := status.FromError(err); ok {
@ -278,13 +364,23 @@ func newClientStream(ctx context.Context, desc *StreamDesc, cc *ClientConn, meth
ctx = rpcConfig.Context
}
mc = &rpcConfig.MethodConfig
onCommit = rpcConfig.OnCommitted
if rpcConfig.OnCommitted != nil {
onCommit = rpcConfig.OnCommitted
// Register an OnFinish CallOption with the OnCommitted callback to
// ensure it is invoked on stream termination, even if the stream
// fails early before committing. Implementations of OnCommitted are
// expected to be idempotent (e.g., guarded by sync.Once), since both
// onCommit and OnFinish may run for a single RPC.
opts = append(opts, OnFinish(func(error) { rpcConfig.OnCommitted() }))
}
if rpcConfig.Interceptor != nil {
rpcInfo.Context = nil
ns := newStream
if interceptor, ok := rpcConfig.Interceptor.(clientInterceptor); ok {
newStream = func(ctx context.Context, filterOpts ...CallOption) (ClientStream, error) {
cs, err := interceptor.NewStream(ctx, rpcInfo, ns, filterOpts...)
newStream = func(ctx context.Context, opts ...CallOption) (ClientStream, error) {
cs, err := interceptor.NewStream(ctx, rpcInfo, ns, opts...)
if err != nil {
return nil, toRPCErr(err)
}
@ -296,7 +392,7 @@ func newClientStream(ctx context.Context, desc *StreamDesc, cc *ClientConn, meth
}
}
return newStream(ctx)
return newStream(ctx, opts...)
}
func newClientStreamWithParams(ctx context.Context, desc *StreamDesc, cc *ClientConn, method string, mc *serviceconfig.MethodConfig, onCommit func(), nameResolutionDelayed bool, opts ...CallOption) (_ ClientStream, err error) {
@ -543,6 +639,9 @@ func (a *csAttempt) newStream() error {
// maintained in it are local to the attempt. When the attempt has to be
// retried, a new instance of csAttempt will be created.
if a.pickResult.Metadata != nil {
if err := imetadata.Validate(a.pickResult.Metadata); err != nil {
return status.Error(codes.Internal, err.Error())
}
// We currently do not have a function it the metadata package which
// merges given metadata with existing metadata in a context. Existing
// function `AppendToOutgoingContext()` takes a variadic argument of key
@ -1044,8 +1143,8 @@ func (cs *clientStream) RecvMsg(m any) error {
binlog.Log(cs.ctx, sm)
}
}
if err != nil || !cs.desc.ServerStreams {
// err != nil or non-server-streaming indicates end of stream.
if err != nil {
// err != nil indicates end of stream.
cs.finish(err)
}
return err
@ -1093,7 +1192,8 @@ func (cs *clientStream) finish(err error) {
}
cs.finished = true
cs.commitAttemptLocked()
if cs.attempt != nil {
attemptCreated := cs.attempt != nil
if attemptCreated {
cs.attempt.finish(err)
// after functions all rely upon having a stream.
if cs.attempt.transportStream != nil {
@ -1131,7 +1231,13 @@ func (cs *clientStream) finish(err error) {
if err == nil {
cs.retryThrottler.successfulRPC()
}
endOfClientStream(cs.cc, err, cs.opts...)
// If no attempt was ever created, stream creation has failed, and the
// cleanup is left to newClientStream, whose deferred cleanup invokes
// endOfClientStream if the call fails. Invoking it here as well would
// run the cleanup twice for the same call.
if attemptCreated {
endOfClientStream(cs.cc, err, cs.opts...)
}
cs.cancel()
}
@ -1145,12 +1251,6 @@ func (a *csAttempt) sendMsg(m any, hdr []byte, payld mem.BufferSlice, dataLength
a.mu.Unlock()
}
if err := a.transportStream.Write(hdr, payld, &transport.WriteOptions{Last: !cs.desc.ClientStreams}); err != nil {
if !cs.desc.ClientStreams {
// For non-client-streaming RPCs, we return nil instead of EOF on error
// because the generated code requires it. finish is not called; RecvMsg()
// will call it with the stream's status independently.
return nil
}
return io.EOF
}
if a.statsHandler != nil {
@ -1218,18 +1318,7 @@ func (a *csAttempt) recvMsg(m any, payInfo *payloadInfo) (err error) {
Length: payInfo.uncompressedBytes.Len(),
})
}
if cs.desc.ServerStreams {
// Subsequent messages should be received by subsequent RecvMsg calls.
return nil
}
// Special handling for non-server-stream rpcs.
// This recv expects EOF or errors, so we don't collect inPayload.
if err := recv(&a.parser, cs.codec, a.transportStream, a.decompressorV0, m, *cs.callInfo.maxReceiveMessageSize, nil, a.decompressorV1, false); err == io.EOF {
return a.transportStream.Status().Err() // non-server streaming Recv returns nil on success
} else if err != nil {
return toRPCErr(err)
}
return status.Error(codes.Internal, "cardinality violation: expected <EOF> for non server-streaming RPCs, but received another message")
return nil
}
func (a *csAttempt) finish(err error) {
@ -1399,7 +1488,7 @@ func newNonRetryClientStream(ctx context.Context, desc *StreamDesc, method strin
}
}()
}
return as, nil
return &clientStreamWrapper{ClientStream: as, desc: desc}, nil
}
type addrConnStream struct {
@ -1497,12 +1586,6 @@ func (as *addrConnStream) SendMsg(m any) (err error) {
}
if err := as.transportStream.Write(hdr, payload, &transport.WriteOptions{Last: !as.desc.ClientStreams}); err != nil {
if !as.desc.ClientStreams {
// For non-client-streaming RPCs, we return nil instead of EOF on error
// because the generated code requires it. finish is not called; RecvMsg()
// will call it with the stream's status independently.
return nil
}
return io.EOF
}
@ -1511,8 +1594,8 @@ func (as *addrConnStream) SendMsg(m any) (err error) {
func (as *addrConnStream) RecvMsg(m any) (err error) {
defer func() {
if err != nil || !as.desc.ServerStreams {
// err != nil or non-server-streaming indicates end of stream.
if err != nil {
// err != nil indicates end of stream.
as.finish(err)
}
}()
@ -1551,20 +1634,7 @@ func (as *addrConnStream) RecvMsg(m any) (err error) {
return toRPCErr(err)
}
as.receivedFirstMsg = true
if as.desc.ServerStreams {
// Subsequent messages should be received by subsequent RecvMsg calls.
return nil
}
// Special handling for non-server-stream rpcs.
// This recv expects EOF or errors, so we don't collect inPayload.
if err := recv(&as.parser, as.codec, as.transportStream, as.decompressorV0, m, *as.callInfo.maxReceiveMessageSize, nil, as.decompressorV1, false); err == io.EOF {
return as.transportStream.Status().Err() // non-server streaming Recv returns nil on success
} else if err != nil {
return toRPCErr(err)
}
return status.Error(codes.Internal, "cardinality violation: expected <EOF> for non server-streaming RPCs, but received another message")
return nil
}
func (as *addrConnStream) finish(err error) {
@ -1675,6 +1745,8 @@ type serverStream struct {
// synchronized.
serverHeaderBinlogged bool
statusWritten atomic.Bool // True if status has been written to the transport.
mu sync.Mutex // protects trInfo.tr after the service handler runs.
}
@ -1723,6 +1795,16 @@ func (ss *serverStream) SetTrailer(md metadata.MD) {
ss.s.SetTrailer(md)
}
// writeStatus sends the status of a stream to the client. It uses an atomic
// CAS to guarantee that the status is written to the transport exactly once,
// even if called concurrently.
func (ss *serverStream) writeStatus(st *status.Status) error {
if !ss.statusWritten.CompareAndSwap(false, true) {
return nil
}
return ss.s.WriteStatus(st)
}
func (ss *serverStream) SendMsg(m any) (err error) {
defer func() {
if ss.trInfo != nil {
@ -1739,7 +1821,7 @@ func (ss *serverStream) SendMsg(m any) (err error) {
}
if err != nil && err != io.EOF {
st, _ := status.FromError(toRPCErr(err))
ss.s.WriteStatus(st)
ss.writeStatus(st)
// Non-user specified status was sent out. This should be an error
// case (as a server side Cancel maybe).
//
@ -1822,7 +1904,7 @@ func (ss *serverStream) RecvMsg(m any) (err error) {
}
if err != nil && err != io.EOF {
st, _ := status.FromError(toRPCErr(err))
ss.s.WriteStatus(st)
ss.writeStatus(st)
// Non-user specified status was sent out. This should be an error
// case (as a server side Cancel maybe).
//

View file

@ -19,4 +19,4 @@
package grpc
// Version is the current grpc version.
const Version = "1.83.2"
const Version = "1.84.0"

2
vendor/modules.txt vendored
View file

@ -1089,7 +1089,7 @@ google.golang.org/genproto/googleapis/api/annotations
# google.golang.org/genproto/googleapis/rpc v0.0.0-20260825221802-da73d73af1c5
## explicit; go 1.25.0
google.golang.org/genproto/googleapis/rpc/status
# google.golang.org/grpc v1.83.2
# google.golang.org/grpc v1.84.0
## explicit; go 1.25.0
google.golang.org/grpc
google.golang.org/grpc/attributes