fix: close stream adaptive gaps and switch notify to stario v0.1.1
- make stream fast path honor adaptive soft payload limits end-to-end - split oversized fast-stream payloads into sequential frames before batching - use adaptive soft cap when encoding stream batch payloads - move timeout-like error detection into production code for adaptive tx - tune notify FrameReader read size explicitly to avoid throughput regression - drop local stario replace and depend on released b612.me/stario v0.1.1
This commit is contained in:
+168
-19
@@ -40,6 +40,8 @@ type modernPSKTransportBundle struct {
|
||||
fastPlainEncode transportFastPlainEncoder
|
||||
}
|
||||
|
||||
var modernPSKPayloadPool sync.Pool
|
||||
|
||||
// ModernPSKOptions configures the modern PSK transport profile.
|
||||
//
|
||||
// The current profile derives a 32-byte transport key with Argon2id and uses
|
||||
@@ -81,6 +83,10 @@ func UseModernPSKClient(c Client, sharedSecret []byte, opts *ModernPSKOptions) e
|
||||
return err
|
||||
}
|
||||
transport := buildModernPSKTransportBundle(aad)
|
||||
runtime, err := newModernPSKCodecRuntime(key, aad)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
c.SetSecretKey(key)
|
||||
c.SetMsgEn(transport.msgEn)
|
||||
c.SetMsgDe(transport.msgDe)
|
||||
@@ -88,6 +94,7 @@ func UseModernPSKClient(c Client, sharedSecret []byte, opts *ModernPSKOptions) e
|
||||
client.fastStreamEncode = transport.fastStreamEncode
|
||||
client.fastBulkEncode = transport.fastBulkEncode
|
||||
client.fastPlainEncode = transport.fastPlainEncode
|
||||
client.modernPSKRuntime = runtime
|
||||
}
|
||||
c.SetSkipExchangeKey(true)
|
||||
return nil
|
||||
@@ -104,6 +111,10 @@ func UseModernPSKServer(s Server, sharedSecret []byte, opts *ModernPSKOptions) e
|
||||
return err
|
||||
}
|
||||
transport := buildModernPSKTransportBundle(aad)
|
||||
runtime, err := newModernPSKCodecRuntime(key, aad)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
s.SetSecretKey(key)
|
||||
s.SetDefaultCommEncode(transport.msgEn)
|
||||
s.SetDefaultCommDecode(transport.msgDe)
|
||||
@@ -111,6 +122,7 @@ func UseModernPSKServer(s Server, sharedSecret []byte, opts *ModernPSKOptions) e
|
||||
server.defaultFastStreamEncode = transport.fastStreamEncode
|
||||
server.defaultFastBulkEncode = transport.fastBulkEncode
|
||||
server.defaultFastPlainEncode = transport.fastPlainEncode
|
||||
server.defaultModernPSKRuntime = runtime
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -127,6 +139,7 @@ func UseLegacySecurityClient(c Client) {
|
||||
client.fastStreamEncode = nil
|
||||
client.fastBulkEncode = nil
|
||||
client.fastPlainEncode = nil
|
||||
client.modernPSKRuntime = nil
|
||||
}
|
||||
c.SetSkipExchangeKey(false)
|
||||
c.SetRsaPubKey(bytes.Clone(defaultRsaPubKey))
|
||||
@@ -144,6 +157,7 @@ func UseLegacySecurityServer(s Server) {
|
||||
server.defaultFastStreamEncode = nil
|
||||
server.defaultFastBulkEncode = nil
|
||||
server.defaultFastPlainEncode = nil
|
||||
server.defaultModernPSKRuntime = nil
|
||||
}
|
||||
s.SetRsaPrivKey(bytes.Clone(defaultRsaKey))
|
||||
}
|
||||
@@ -185,14 +199,14 @@ func buildModernPSKCodecs(aad []byte) (func([]byte, []byte) []byte, func([]byte,
|
||||
|
||||
func buildModernPSKTransportBundle(aad []byte) modernPSKTransportBundle {
|
||||
aadCopy := bytes.Clone(aad)
|
||||
cache := &modernPSKCodecCache{}
|
||||
cache := newModernPSKCodecCache(aadCopy)
|
||||
msgEn := func(key []byte, plain []byte) []byte {
|
||||
runtime, err := cache.runtimeForKey(key)
|
||||
if err != nil {
|
||||
log.Print(err)
|
||||
return nil
|
||||
}
|
||||
out, err := runtime.sealPlainPayload(aadCopy, plain)
|
||||
out, err := runtime.sealPlainPayload(plain)
|
||||
if err != nil {
|
||||
log.Print(err)
|
||||
return nil
|
||||
@@ -214,9 +228,7 @@ func buildModernPSKTransportBundle(aad []byte) modernPSKTransportBundle {
|
||||
log.Print(err)
|
||||
return nil
|
||||
}
|
||||
nonce := encrypted[len(modernPSKMagic):headerLen]
|
||||
ciphertext := encrypted[headerLen:]
|
||||
plain, err := runtime.aead.Open(make([]byte, 0, len(ciphertext)), nonce, ciphertext, aadCopy)
|
||||
plain, err := runtime.openPayload(encrypted)
|
||||
if err != nil {
|
||||
log.Print(err)
|
||||
return nil
|
||||
@@ -228,21 +240,21 @@ func buildModernPSKTransportBundle(aad []byte) modernPSKTransportBundle {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return runtime.sealStreamFastPayload(aadCopy, dataID, seq, payload)
|
||||
return runtime.sealStreamFastPayload(dataID, seq, payload)
|
||||
}
|
||||
fastBulkEncode := func(key []byte, dataID uint64, seq uint64, payload []byte) ([]byte, error) {
|
||||
runtime, err := cache.runtimeForKey(key)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return runtime.sealBulkFastPayload(aadCopy, dataID, seq, payload)
|
||||
return runtime.sealBulkFastPayload(dataID, seq, payload)
|
||||
}
|
||||
fastPlainEncode := func(key []byte, plainLen int, fill func([]byte) error) ([]byte, error) {
|
||||
runtime, err := cache.runtimeForKey(key)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return runtime.sealFilledPayload(aadCopy, plainLen, fill)
|
||||
return runtime.sealFilledPayload(plainLen, fill)
|
||||
}
|
||||
return modernPSKTransportBundle{
|
||||
msgEn: msgEn,
|
||||
@@ -269,16 +281,23 @@ func (s *ServerCommon) validateSecurityConfiguration() error {
|
||||
|
||||
type modernPSKCodecCache struct {
|
||||
mu sync.Mutex
|
||||
aad []byte
|
||||
key []byte
|
||||
runtime *modernPSKCodecRuntime
|
||||
}
|
||||
|
||||
type modernPSKCodecRuntime struct {
|
||||
aead cipher.AEAD
|
||||
key []byte
|
||||
aad []byte
|
||||
prefix [modernPSKNonceSize - 8]byte
|
||||
seq atomic.Uint64
|
||||
}
|
||||
|
||||
func newModernPSKCodecCache(aad []byte) *modernPSKCodecCache {
|
||||
return &modernPSKCodecCache{aad: bytes.Clone(aad)}
|
||||
}
|
||||
|
||||
func (c *modernPSKCodecCache) runtimeForKey(key []byte) (*modernPSKCodecRuntime, error) {
|
||||
if c == nil {
|
||||
return nil, errModernPSKSecretEmpty
|
||||
@@ -288,7 +307,7 @@ func (c *modernPSKCodecCache) runtimeForKey(key []byte) (*modernPSKCodecRuntime,
|
||||
if c.runtime != nil && bytes.Equal(c.key, key) {
|
||||
return c.runtime, nil
|
||||
}
|
||||
runtime, err := newModernPSKCodecRuntime(key)
|
||||
runtime, err := newModernPSKCodecRuntime(key, c.aad)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -297,7 +316,7 @@ func (c *modernPSKCodecCache) runtimeForKey(key []byte) (*modernPSKCodecRuntime,
|
||||
return runtime, nil
|
||||
}
|
||||
|
||||
func newModernPSKCodecRuntime(key []byte) (*modernPSKCodecRuntime, error) {
|
||||
func newModernPSKCodecRuntime(key []byte, aad []byte) (*modernPSKCodecRuntime, error) {
|
||||
if len(key) == 0 {
|
||||
return nil, errModernPSKSecretEmpty
|
||||
}
|
||||
@@ -311,6 +330,8 @@ func newModernPSKCodecRuntime(key []byte) (*modernPSKCodecRuntime, error) {
|
||||
}
|
||||
runtime := &modernPSKCodecRuntime{
|
||||
aead: aead,
|
||||
key: bytes.Clone(key),
|
||||
aad: bytes.Clone(aad),
|
||||
}
|
||||
if _, err := cryptorand.Read(runtime.prefix[:]); err != nil {
|
||||
return nil, err
|
||||
@@ -318,6 +339,13 @@ func newModernPSKCodecRuntime(key []byte) (*modernPSKCodecRuntime, error) {
|
||||
return runtime, nil
|
||||
}
|
||||
|
||||
func (r *modernPSKCodecRuntime) fork() (*modernPSKCodecRuntime, error) {
|
||||
if r == nil {
|
||||
return nil, errModernPSKSecretEmpty
|
||||
}
|
||||
return newModernPSKCodecRuntime(r.key, r.aad)
|
||||
}
|
||||
|
||||
func (r *modernPSKCodecRuntime) nextNonce() [modernPSKNonceSize]byte {
|
||||
var nonce [modernPSKNonceSize]byte
|
||||
if r == nil {
|
||||
@@ -328,8 +356,8 @@ func (r *modernPSKCodecRuntime) nextNonce() [modernPSKNonceSize]byte {
|
||||
return nonce
|
||||
}
|
||||
|
||||
func (r *modernPSKCodecRuntime) sealStreamFastPayload(aad []byte, dataID uint64, seq uint64, payload []byte) ([]byte, error) {
|
||||
return r.sealFilledPayload(aad, streamFastPayloadHeaderLen+len(payload), func(frame []byte) error {
|
||||
func (r *modernPSKCodecRuntime) sealStreamFastPayload(dataID uint64, seq uint64, payload []byte) ([]byte, error) {
|
||||
return r.sealFilledPayload(streamFastPayloadHeaderLen+len(payload), func(frame []byte) error {
|
||||
if err := encodeStreamFastDataFrameHeader(frame, dataID, seq, len(payload)); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -338,11 +366,11 @@ func (r *modernPSKCodecRuntime) sealStreamFastPayload(aad []byte, dataID uint64,
|
||||
})
|
||||
}
|
||||
|
||||
func (r *modernPSKCodecRuntime) sealBulkFastPayload(aad []byte, dataID uint64, seq uint64, payload []byte) ([]byte, error) {
|
||||
func (r *modernPSKCodecRuntime) sealBulkFastPayload(dataID uint64, seq uint64, payload []byte) ([]byte, error) {
|
||||
if r == nil {
|
||||
return nil, errTransportPayloadEncryptFailed
|
||||
}
|
||||
return r.sealFilledPayload(aad, bulkFastPayloadHeaderLen+len(payload), func(frame []byte) error {
|
||||
return r.sealFilledPayload(bulkFastPayloadHeaderLen+len(payload), func(frame []byte) error {
|
||||
if err := encodeBulkFastDataFrameHeader(frame, dataID, seq, len(payload)); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -351,14 +379,14 @@ func (r *modernPSKCodecRuntime) sealBulkFastPayload(aad []byte, dataID uint64, s
|
||||
})
|
||||
}
|
||||
|
||||
func (r *modernPSKCodecRuntime) sealPlainPayload(aad []byte, plain []byte) ([]byte, error) {
|
||||
return r.sealFilledPayload(aad, len(plain), func(dst []byte) error {
|
||||
func (r *modernPSKCodecRuntime) sealPlainPayload(plain []byte) ([]byte, error) {
|
||||
return r.sealFilledPayload(len(plain), func(dst []byte) error {
|
||||
copy(dst, plain)
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func (r *modernPSKCodecRuntime) sealFilledPayload(aad []byte, plainLen int, fill func([]byte) error) ([]byte, error) {
|
||||
func (r *modernPSKCodecRuntime) sealFilledPayload(plainLen int, fill func([]byte) error) ([]byte, error) {
|
||||
if r == nil {
|
||||
return nil, errTransportPayloadEncryptFailed
|
||||
}
|
||||
@@ -368,6 +396,35 @@ func (r *modernPSKCodecRuntime) sealFilledPayload(aad []byte, plainLen int, fill
|
||||
nonce := r.nextNonce()
|
||||
headerLen := len(modernPSKMagic) + modernPSKNonceSize
|
||||
out := make([]byte, headerLen+plainLen+r.aead.Overhead())
|
||||
sealed, err := r.sealInto(out, headerLen, nonce, plainLen, fill)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out[:headerLen+len(sealed)], nil
|
||||
}
|
||||
|
||||
func (r *modernPSKCodecRuntime) sealFilledPayloadPooled(plainLen int, fill func([]byte) error) ([]byte, func(), error) {
|
||||
if r == nil {
|
||||
return nil, nil, errTransportPayloadEncryptFailed
|
||||
}
|
||||
if plainLen < 0 {
|
||||
return nil, nil, errTransportPayloadEncryptFailed
|
||||
}
|
||||
nonce := r.nextNonce()
|
||||
headerLen := len(modernPSKMagic) + modernPSKNonceSize
|
||||
totalLen := headerLen + plainLen + r.aead.Overhead()
|
||||
out := getModernPSKPayloadBuffer(totalLen)
|
||||
sealed, err := r.sealInto(out, headerLen, nonce, plainLen, fill)
|
||||
if err != nil {
|
||||
putModernPSKPayloadBuffer(out)
|
||||
return nil, nil, err
|
||||
}
|
||||
return out[:headerLen+len(sealed)], func() {
|
||||
putModernPSKPayloadBuffer(out)
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (r *modernPSKCodecRuntime) sealInto(out []byte, headerLen int, nonce [modernPSKNonceSize]byte, plainLen int, fill func([]byte) error) ([]byte, error) {
|
||||
copy(out[:len(modernPSKMagic)], modernPSKMagic)
|
||||
copy(out[len(modernPSKMagic):headerLen], nonce[:])
|
||||
frame := out[headerLen : headerLen+plainLen]
|
||||
@@ -376,6 +433,98 @@ func (r *modernPSKCodecRuntime) sealFilledPayload(aad []byte, plainLen int, fill
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
sealed := r.aead.Seal(frame[:0], nonce[:], frame, aad)
|
||||
return out[:headerLen+len(sealed)], nil
|
||||
return r.aead.Seal(frame[:0], nonce[:], frame, r.aad), nil
|
||||
}
|
||||
|
||||
func (r *modernPSKCodecRuntime) openPayload(encrypted []byte) ([]byte, error) {
|
||||
if r == nil {
|
||||
return nil, errTransportPayloadDecryptFailed
|
||||
}
|
||||
headerLen := len(modernPSKMagic) + modernPSKNonceSize
|
||||
if len(encrypted) < headerLen {
|
||||
return nil, errModernPSKPayload
|
||||
}
|
||||
if !bytes.Equal(encrypted[:len(modernPSKMagic)], modernPSKMagic) {
|
||||
return nil, errModernPSKPayload
|
||||
}
|
||||
nonce := encrypted[len(modernPSKMagic):headerLen]
|
||||
ciphertext := encrypted[headerLen:]
|
||||
return r.aead.Open(make([]byte, 0, len(ciphertext)), nonce, ciphertext, r.aad)
|
||||
}
|
||||
|
||||
func (r *modernPSKCodecRuntime) openPayloadPooled(encrypted []byte, release func()) ([]byte, func(), error) {
|
||||
if r == nil {
|
||||
if release != nil {
|
||||
release()
|
||||
}
|
||||
return nil, nil, errTransportPayloadDecryptFailed
|
||||
}
|
||||
headerLen := len(modernPSKMagic) + modernPSKNonceSize
|
||||
if len(encrypted) < headerLen {
|
||||
if release != nil {
|
||||
release()
|
||||
}
|
||||
return nil, nil, errModernPSKPayload
|
||||
}
|
||||
if !bytes.Equal(encrypted[:len(modernPSKMagic)], modernPSKMagic) {
|
||||
if release != nil {
|
||||
release()
|
||||
}
|
||||
return nil, nil, errModernPSKPayload
|
||||
}
|
||||
nonce := encrypted[len(modernPSKMagic):headerLen]
|
||||
ciphertext := encrypted[headerLen:]
|
||||
plain, err := r.aead.Open(ciphertext[:0], nonce, ciphertext, r.aad)
|
||||
if err != nil {
|
||||
if release != nil {
|
||||
release()
|
||||
}
|
||||
return nil, nil, err
|
||||
}
|
||||
return plain, release, nil
|
||||
}
|
||||
|
||||
func (r *modernPSKCodecRuntime) openPayloadOwnedPooled(encrypted []byte) ([]byte, func(), error) {
|
||||
if r == nil {
|
||||
return nil, nil, errTransportPayloadDecryptFailed
|
||||
}
|
||||
headerLen := len(modernPSKMagic) + modernPSKNonceSize
|
||||
if len(encrypted) < headerLen {
|
||||
return nil, nil, errModernPSKPayload
|
||||
}
|
||||
if !bytes.Equal(encrypted[:len(modernPSKMagic)], modernPSKMagic) {
|
||||
return nil, nil, errModernPSKPayload
|
||||
}
|
||||
nonce := encrypted[len(modernPSKMagic):headerLen]
|
||||
ciphertext := encrypted[headerLen:]
|
||||
plainLen := len(ciphertext) - r.aead.Overhead()
|
||||
if plainLen < 0 {
|
||||
return nil, nil, errModernPSKPayload
|
||||
}
|
||||
out := getModernPSKPayloadBuffer(plainLen)
|
||||
plain, err := r.aead.Open(out[:0], nonce, ciphertext, r.aad)
|
||||
if err != nil {
|
||||
putModernPSKPayloadBuffer(out)
|
||||
return nil, nil, err
|
||||
}
|
||||
return plain, func() {
|
||||
putModernPSKPayloadBuffer(out)
|
||||
}, nil
|
||||
}
|
||||
|
||||
func getModernPSKPayloadBuffer(size int) []byte {
|
||||
if size <= 0 {
|
||||
return nil
|
||||
}
|
||||
if pooled, ok := modernPSKPayloadPool.Get().([]byte); ok && cap(pooled) >= size {
|
||||
return pooled[:size]
|
||||
}
|
||||
return make([]byte, size)
|
||||
}
|
||||
|
||||
func putModernPSKPayloadBuffer(buf []byte) {
|
||||
if cap(buf) == 0 || cap(buf) > 32*1024*1024 {
|
||||
return
|
||||
}
|
||||
modernPSKPayloadPool.Put(buf[:0])
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user