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 }