Files
notify/stream_runtime.go
T

485 lines
11 KiB
Go
Raw Permalink Normal View History

package notify
import (
"context"
"fmt"
"strings"
"sync"
"sync/atomic"
)
type streamRuntime struct {
rolePrefix string
seq atomic.Uint64
dataSeq uint64
peerDataSeq uint64
dataStart uint64
dataStep uint64
mu sync.RWMutex
handler func(StreamAcceptInfo) error
streams map[string]*streamHandle
data map[string]map[uint64]*streamHandle
reserved map[string]map[uint64]struct{}
cfg streamConfig
flow *streamFlowController
}
func newStreamRuntime(rolePrefix string) *streamRuntime {
cfg := defaultStreamConfig()
dataStart, dataStep := uint64(1), uint64(1)
if rolePrefix == "cstrm" {
dataStep = 2
} else if rolePrefix == "sstrm" {
dataStart = 2
dataStep = 2
}
return &streamRuntime{
rolePrefix: rolePrefix,
dataStart: dataStart,
dataStep: dataStep,
streams: make(map[string]*streamHandle),
data: make(map[string]map[uint64]*streamHandle),
reserved: make(map[string]map[uint64]struct{}),
cfg: cfg,
flow: newStreamFlowController(cfg),
}
}
func (r *streamRuntime) nextID() string {
if r == nil {
return ""
}
return fmt.Sprintf("%s-%d", r.rolePrefix, r.seq.Add(1))
}
func (r *streamRuntime) nextDataID() uint64 {
if r == nil {
return 0
}
r.mu.Lock()
defer r.mu.Unlock()
id, _ := r.nextDataIDLocked(defaultFileScope, false)
return id
}
func (r *streamRuntime) reserveDataID(scope string) (uint64, error) {
if r == nil {
return 0, errStreamRuntimeNil
}
scope = normalizeFileScope(scope)
r.mu.Lock()
defer r.mu.Unlock()
id, err := r.nextDataIDLocked(scope, true)
return id, err
}
func (r *streamRuntime) releaseDataID(scope string, dataID uint64) {
if r == nil || dataID == 0 {
return
}
scope = normalizeFileScope(scope)
r.mu.Lock()
defer r.mu.Unlock()
if reserved := r.reserved[scope]; reserved != nil {
delete(reserved, dataID)
if len(reserved) == 0 {
delete(r.reserved, scope)
}
}
}
func (r *streamRuntime) nextDataIDLocked(scope string, reserve bool) (uint64, error) {
if r == nil {
return 0, errStreamRuntimeNil
}
for {
candidate, ok := r.nextDataCandidateLocked(false)
if !ok {
return 0, errStreamDataIDExhausted
}
if r.dataInUseLocked(scope, candidate) {
continue
}
if reserve {
reserved := r.reserved[scope]
if reserved == nil {
reserved = make(map[uint64]struct{})
r.reserved[scope] = reserved
}
reserved[candidate] = struct{}{}
}
return candidate, nil
}
}
func (r *streamRuntime) nextPeerDataIDLocked(scope string) (uint64, error) {
if r == nil {
return 0, errStreamRuntimeNil
}
for {
candidate, ok := r.nextDataCandidateLocked(true)
if !ok {
return 0, errStreamDataIDExhausted
}
if r.dataInUseLocked(scope, candidate) {
continue
}
return candidate, nil
}
}
func (r *streamRuntime) nextDataCandidateLocked(peer bool) (uint64, bool) {
seq := &r.dataSeq
start, step := r.dataStart, r.dataStep
if peer && step == 2 {
seq = &r.peerDataSeq
start = 3 - r.dataStart
}
if step == 0 {
step = 1
}
if *seq == 0 {
if start == 0 {
return 0, false
}
*seq = start
return start, true
}
if *seq > ^uint64(0)-step {
return 0, false
}
candidate := *seq + step
if step == 2 && candidate%2 != start%2 {
if candidate == ^uint64(0) {
return 0, false
}
candidate++
}
*seq = candidate
return candidate, true
}
func (r *streamRuntime) dataInUseLocked(scope string, dataID uint64) bool {
if dataID == 0 {
return true
}
if dataScope := r.data[scope]; dataScope != nil {
if _, ok := dataScope[dataID]; ok {
return true
}
}
if reserved := r.reserved[scope]; reserved != nil {
_, ok := reserved[dataID]
return ok
}
return false
}
func (r *streamRuntime) setHandler(fn func(StreamAcceptInfo) error) {
if r == nil {
return
}
r.mu.Lock()
defer r.mu.Unlock()
r.handler = fn
}
func (r *streamRuntime) handlerSnapshot() func(StreamAcceptInfo) error {
if r == nil {
return nil
}
r.mu.RLock()
defer r.mu.RUnlock()
return r.handler
}
func (r *streamRuntime) register(scope string, stream *streamHandle) error {
return r.registerWithDirection(scope, stream, false, false)
}
func (r *streamRuntime) registerInbound(scope string, stream *streamHandle) error {
return r.registerWithDirection(scope, stream, true, false)
}
func (r *streamRuntime) registerReserved(scope string, stream *streamHandle) error {
return r.registerWithDirection(scope, stream, false, true)
}
func (r *streamRuntime) registerWithDirection(scope string, stream *streamHandle, inbound bool, consumeReservation bool) error {
if r == nil {
return errStreamRuntimeNil
}
if stream == nil || stream.id == "" {
return errStreamIDEmpty
}
scope = normalizeFileScope(scope)
key := streamRuntimeKey(scope, stream.id)
r.mu.Lock()
defer r.mu.Unlock()
if _, ok := r.streams[key]; ok {
return errStreamAlreadyExists
}
if stream.dataID == 0 {
var err error
if inbound {
stream.dataID, err = r.nextPeerDataIDLocked(scope)
} else {
stream.dataID, err = r.nextDataIDLocked(scope, false)
}
if err != nil {
return err
}
}
if stream.dataID != 0 {
dataScope := r.data[scope]
if dataScope == nil {
dataScope = make(map[uint64]*streamHandle)
r.data[scope] = dataScope
}
if _, ok := dataScope[stream.dataID]; ok {
return errStreamAlreadyExists
}
if reserved := r.reserved[scope]; reserved != nil {
if _, ok := reserved[stream.dataID]; ok {
if !consumeReservation {
return errStreamAlreadyExists
}
delete(reserved, stream.dataID)
if len(reserved) == 0 {
delete(r.reserved, scope)
}
}
}
dataScope[stream.dataID] = stream
}
r.streams[key] = stream
return nil
}
// adopt transfers ownership of a newly-created stream to the runtime. A
// failed registration is terminal so its child context cannot remain attached
// to the session after the caller drops the handle.
func (r *streamRuntime) adopt(scope string, stream *streamHandle) error {
err := r.register(scope, stream)
if err != nil && stream != nil {
stream.markReset(err)
}
return err
}
func (r *streamRuntime) adoptInbound(scope string, stream *streamHandle) error {
err := r.registerInbound(scope, stream)
if err != nil && stream != nil {
stream.markReset(err)
}
return err
}
func (r *streamRuntime) adoptReserved(scope string, stream *streamHandle) error {
err := r.registerReserved(scope, stream)
if err != nil && stream != nil {
stream.markReset(err)
}
return err
}
func (r *streamRuntime) lookup(scope string, streamID string) (*streamHandle, bool) {
if r == nil || streamID == "" {
return nil, false
}
key := streamRuntimeKey(scope, streamID)
r.mu.RLock()
defer r.mu.RUnlock()
stream, ok := r.streams[key]
return stream, ok
}
func (r *streamRuntime) lookupByDataID(scope string, dataID uint64) (*streamHandle, bool) {
if r == nil || dataID == 0 {
return nil, false
}
scope = normalizeFileScope(scope)
r.mu.RLock()
defer r.mu.RUnlock()
dataScope := r.data[scope]
if dataScope == nil {
return nil, false
}
stream, ok := dataScope[dataID]
return stream, ok
}
func (r *streamRuntime) lookupControl(scope string, streamID string, dataID uint64) (*streamHandle, bool) {
if r == nil {
return nil, false
}
scope = normalizeFileScope(scope)
r.mu.RLock()
defer r.mu.RUnlock()
if streamID != "" {
stream, ok := r.streams[streamRuntimeKey(scope, streamID)]
if !ok || stream == nil {
return nil, false
}
if dataID != 0 && stream.dataID != dataID {
return nil, false
}
return stream, true
}
if dataID == 0 {
return nil, false
}
stream := r.data[scope][dataID]
return stream, stream != nil
}
func (r *streamRuntime) remove(scope string, expected *streamHandle) {
if r == nil || expected == nil || expected.id == "" {
return
}
scope = normalizeFileScope(scope)
key := streamRuntimeKey(scope, expected.id)
r.mu.Lock()
defer r.mu.Unlock()
stream := r.streams[key]
if stream != expected {
return
}
if stream.dataID != 0 {
if dataScope := r.data[scope]; dataScope != nil {
if dataScope[stream.dataID] == stream {
delete(dataScope, stream.dataID)
}
if len(dataScope) == 0 {
delete(r.data, scope)
}
}
}
delete(r.streams, key)
}
func (r *streamRuntime) acquireOutbound(ctx context.Context, size int) (func(), error) {
if r == nil || r.flow == nil {
return func() {}, nil
}
return r.flow.acquire(ctx, size)
}
func (r *streamRuntime) tryAcquireOutbound(size int) bool {
if r == nil || r.flow == nil {
return true
}
return r.flow.tryAcquire(size)
}
func (r *streamRuntime) releaseOutbound(size int) {
if r == nil || r.flow == nil {
return
}
r.flow.release(size)
}
func (r *streamRuntime) snapshots() []StreamSnapshot {
if r == nil {
return nil
}
r.mu.RLock()
snapshots := make([]StreamSnapshot, 0, len(r.streams))
for _, stream := range r.streams {
if stream == nil {
continue
}
snapshots = append(snapshots, stream.snapshot())
}
r.mu.RUnlock()
sortStreamSnapshots(snapshots)
return snapshots
}
func (r *streamRuntime) closeAll(err error) {
r.closeMatching(func(string) bool { return true }, err)
}
func (r *streamRuntime) closeScope(scope string, err error) {
scope = normalizeFileScope(scope)
r.closeMatching(func(key string) bool {
return strings.HasPrefix(key, scope+"\x00")
}, err)
}
func (r *streamRuntime) closeClientRoute(route clientSessionRoute, err error) {
if r == nil {
return
}
if !r.mu.TryRLock() {
go r.closeClientRouteBlocking(route, err)
return
}
streams := r.collectClientRouteLocked(route)
r.mu.RUnlock()
r.resetClientRouteHandles(streams, err)
}
func (r *streamRuntime) closeClientRouteBlocking(route clientSessionRoute, err error) {
r.mu.RLock()
streams := r.collectClientRouteLocked(route)
r.mu.RUnlock()
r.resetClientRouteHandles(streams, err)
}
func (r *streamRuntime) collectClientRouteLocked(route clientSessionRoute) []*streamHandle {
streams := make([]*streamHandle, 0)
for _, stream := range r.streams {
if stream == nil || !sameClientSessionRoute(stream.clientRoute, route) {
continue
}
streams = append(streams, stream)
}
return streams
}
func (r *streamRuntime) resetClientRouteHandles(streams []*streamHandle, err error) {
resetErr := streamRuntimeCloseError(err)
for _, stream := range streams {
stream.markReset(resetErr)
}
}
func (r *streamRuntime) closeMatching(match func(string) bool, err error) {
if r == nil || match == nil {
return
}
resetErr := streamRuntimeCloseError(err)
r.mu.RLock()
streams := make([]*streamHandle, 0, len(r.streams))
for key, stream := range r.streams {
if stream == nil || !match(key) {
continue
}
streams = append(streams, stream)
}
r.mu.RUnlock()
for _, stream := range streams {
stream.markReset(resetErr)
}
}
func streamRuntimeKey(scope string, streamID string) string {
return normalizeFileScope(scope) + "\x00" + streamID
}
func (c *ClientCommon) getStreamRuntime() *streamRuntime {
if c == nil {
return nil
}
return c.streamRuntime
}
func (s *ServerCommon) getStreamRuntime() *streamRuntime {
if s == nil {
return nil
}
return s.streamRuntime
}