This is an automated email from the ASF dual-hosted git repository.
sruehl pushed a commit to branch develop
in repository https://gitbox.apache.org/repos/asf/plc4x.git
The following commit(s) were added to refs/heads/develop by this push:
new 3a589b28f9 Connection lease management and modbus improvements (#2570)
3a589b28f9 is described below
commit 3a589b28f93e936566e2b0445615efd29ff2220b
Author: Shaun <[email protected]>
AuthorDate: Thu May 28 01:24:25 2026 -0700
Connection lease management and modbus improvements (#2570)
* fix: harden Modbus Receive against truncated/extended frames
- retry parsing with full buffered data when initial MBAP length is wrong
- treat EOF as incomplete data and avoid dropping partial frames
- keep discarding truly unparsable packets with diagnostic logging
* fix: watch for and discard trailing CRCs from misbehaving gateways
fix: attempt to resynchronize the read stream if desynchronized
* fix: additional sanity checks on the MBAP
fix: io.EOF is a trap, we checked fragmentation above
* fix: deal with TCP keep-alive padding that leaks from the kernel
* refactor: attempting to simply logic while still covering desync handling
* refactor: more robust consistency checks
* fix: Final consistency check case should discard all available bytes
* refactor: reduce log spam
* fix: handleDesync more robustly handles the padding leak issue
* fix: keep the last 5 bytes to avoid breaking fragmentation
* feat: connections can now Invalidate() to mark irrecoverably failed
* feat: propagate transport error classification through codecs
* add TransportErrorKind API with handler plumbing in transports
* wire DefaultCodec/Connection and protocol codecs to classify and react to
fatal errors
* feat(transports): new TransportError type for consistent transport error
propagation
... and easier access
* fix: export TransportErrorKind values and add comments
* fix: few broken behaviors in prior change around io.EOF
* fix(transports): use shared ErrorIs helper for typed nil handling
* fix: another nil != typed_nil path
* fix: more nil != typed_nil
* fix: just Trace level log on double Close()
---
plc4go/go.mod | 1 +
plc4go/go.sum | 2 +
plc4go/internal/ads/DiscoveryMessageCodec.go | 40 +++++-
plc4go/internal/ads/MessageCodec.go | 47 ++++++-
plc4go/internal/bacnetip/MessageCodec_test.go | 4 +
plc4go/internal/cbus/MessageCodec.go | 11 ++
plc4go/internal/knxnetip/Connection.go | 24 ++++
plc4go/internal/modbus/Connection.go | 7 +-
plc4go/internal/modbus/MessageCodec.go | 57 ++++++--
plc4go/internal/simulated/Connection.go | 25 +++-
plc4go/internal/simulated/Connection_test.go | 54 ++++++++
plc4go/pkg/api/PlcConnection.go | 3 +
plc4go/pkg/api/cache/plcConnectionLease.go | 77 +++++++----
plc4go/pkg/api/cache/plcConnectionLease_test.go | 117 +++++++++++++++++
plc4go/pkg/api/transports/transports.go | 28 +++-
plc4go/spi/TransportErrorHandler.go | 9 ++
plc4go/spi/default/DefaultCodec.go | 109 ++++++++++++++-
plc4go/spi/default/DefaultConnection.go | 64 ++++++++-
plc4go/spi/default/mocks_test.go | 131 ++++++++++++++++++
plc4go/spi/transports/TransportInstance.go | 3 +
plc4go/spi/transports/errors.go | 153 ++++++++++++++++++++++
plc4go/spi/transports/mocks_test.go | 91 +++++++++++++
plc4go/spi/transports/pcap/TransportInstance.go | 13 ++
plc4go/spi/transports/serial/TransportInstance.go | 23 ++++
plc4go/spi/transports/tcp/TransportInstance.go | 40 ++++++
plc4go/spi/transports/test/TransportInstance.go | 14 ++
plc4go/spi/transports/udp/TransportInstance.go | 37 ++++++
27 files changed, 1132 insertions(+), 52 deletions(-)
diff --git a/plc4go/go.mod b/plc4go/go.mod
index c989b14ae0..ac3de759d8 100644
--- a/plc4go/go.mod
+++ b/plc4go/go.mod
@@ -29,6 +29,7 @@ require (
github.com/google/uuid v1.6.0
github.com/gopacket/gopacket v1.6.0
github.com/jacobsa/go-serial v0.0.0-20180131005756-15cf729a72d4
+ github.com/pkg/errors v0.9.1
github.com/rs/zerolog v1.35.1
github.com/stretchr/testify v1.11.1
golang.org/x/net v0.55.0
diff --git a/plc4go/go.sum b/plc4go/go.sum
index 6d0be1e259..fd3740d6bc 100644
--- a/plc4go/go.sum
+++ b/plc4go/go.sum
@@ -25,6 +25,8 @@ github.com/mattn/go-colorable v0.1.14/go.mod
h1:6LmQG8QLFO4G5z1gPvYEzlUgJ2wF+stg
github.com/mattn/go-isatty v0.0.22
h1:j8l17JJ9i6VGPUFUYoTUKPSgKe/83EYU2zBC7YNKMw4=
github.com/mattn/go-isatty v0.0.22/go.mod
h1:ZXfXG4SQHsB/w3ZeOYbR0PrPwLy+n6xiMrJlRFqopa4=
github.com/pkg/diff v0.0.0-20210226163009-20ebb0f2a09e/go.mod
h1:pJLUxLENpZxwdsKMEsNbx1VGcRFpLqf3715MtcvvzbA=
+github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4=
+github.com/pkg/errors v0.9.1/go.mod
h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0=
github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2
h1:Jamvg5psRIccs7FGNTlIRMkT8wgtp5eCXdBlqhYGL6U=
github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2/go.mod
h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/rogpeppe/go-internal v1.9.0/go.mod
h1:WtVeX8xhTBvf0smdhujwtBcq4Qrzq/fJaraNFVN+nFs=
diff --git a/plc4go/internal/ads/DiscoveryMessageCodec.go
b/plc4go/internal/ads/DiscoveryMessageCodec.go
index e144f94f67..6e611d4688 100644
--- a/plc4go/internal/ads/DiscoveryMessageCodec.go
+++ b/plc4go/internal/ads/DiscoveryMessageCodec.go
@@ -21,12 +21,13 @@ package ads
import (
"context"
+ "io"
"github.com/rs/zerolog"
"github.com/apache/plc4x/plc4go/protocols/ads/discovery/readwrite/model"
"github.com/apache/plc4x/plc4go/spi"
- "github.com/apache/plc4x/plc4go/spi/default"
+ _default "github.com/apache/plc4x/plc4go/spi/default"
"github.com/apache/plc4x/plc4go/spi/errors"
"github.com/apache/plc4x/plc4go/spi/options"
"github.com/apache/plc4x/plc4go/spi/transports"
@@ -54,6 +55,31 @@ func (m *DiscoveryMessageCodec) GetCodec() spi.MessageCodec {
return m
}
+func (m *DiscoveryMessageCodec) classifyTransportError(err error)
transports.TransportErrorKind {
+ if err == nil {
+ return transports.TransportErrorUnknown
+ }
+ if transport := m.GetTransportInstance(); transport != nil {
+ return transport.ClassifyError(err)
+ }
+ return transports.TransportErrorUnknown
+}
+
+func (m *DiscoveryMessageCodec) isFatalTransportError(err error) bool {
+ if err == nil || transports.ErrorIs(err, io.EOF) {
+ return false
+ }
+ return m.classifyTransportError(err) == transports.TransportErrorFatal
+}
+
+func (m *DiscoveryMessageCodec) wrapFatalTransportError(err error, msg string)
error {
+ if err == nil {
+ return nil
+ }
+ // TODO: Any additional context?
+ return transports.NewTransportError(transports.TransportErrorFatal,
errors.Wrap(err, msg))
+}
+
func (m *DiscoveryMessageCodec) Send(ctx context.Context, interactionInfo
string, message spi.Message) error {
m.log.Trace().Str("interactionInfo", interactionInfo).Msg("Sending
message")
// Cast the message to the correct type of struct
@@ -79,7 +105,9 @@ func (m *DiscoveryMessageCodec) Receive(ctx context.Context)
(spi.Message, error
data, err := m.GetTransportInstance().PeekReadableBytes(ctx, 6)
if err != nil {
m.log.Warn().Err(err).Msg("error peeking")
- // TODO: Possibly clean up ...
+ if m.isFatalTransportError(err) {
+ return nil, m.wrapFatalTransportError(err,
"error peeking header")
+ }
return nil, nil
}
// Get the size of the entire packet little endian plus size of
header
@@ -90,7 +118,10 @@ func (m *DiscoveryMessageCodec) Receive(ctx
context.Context) (spi.Message, error
}
data, err = m.GetTransportInstance().Read(ctx, packetSize)
if err != nil {
- // TODO: Possibly clean up ...
+ m.log.Warn().Err(err).Msg("error reading packet data")
+ if m.isFatalTransportError(err) {
+ return nil, m.wrapFatalTransportError(err,
"error reading packet data")
+ }
return nil, nil
}
ctxForModel := options.GetLoggerContextForModel(ctx, m.log,
options.WithPassLoggerToModel(m.passLogToModel))
@@ -103,6 +134,9 @@ func (m *DiscoveryMessageCodec) Receive(ctx
context.Context) (spi.Message, error
return tcpPacket, nil
} else if err != nil {
m.log.Warn().Err(err).Msg("Got error reading")
+ if m.isFatalTransportError(err) {
+ return nil, m.wrapFatalTransportError(err, "error
getting readable bytes")
+ }
return nil, nil
}
// TODO: maybe we return here a not enough error error
diff --git a/plc4go/internal/ads/MessageCodec.go
b/plc4go/internal/ads/MessageCodec.go
index 5901494771..c4697d7db1 100644
--- a/plc4go/internal/ads/MessageCodec.go
+++ b/plc4go/internal/ads/MessageCodec.go
@@ -22,12 +22,13 @@ package ads
import (
"context"
"encoding/binary"
+ "io"
"github.com/rs/zerolog"
"github.com/apache/plc4x/plc4go/protocols/ads/readwrite/model"
"github.com/apache/plc4x/plc4go/spi"
- "github.com/apache/plc4x/plc4go/spi/default"
+ _default "github.com/apache/plc4x/plc4go/spi/default"
"github.com/apache/plc4x/plc4go/spi/errors"
"github.com/apache/plc4x/plc4go/spi/options"
"github.com/apache/plc4x/plc4go/spi/transports"
@@ -71,6 +72,31 @@ func (m *MessageCodec) GetCodec() spi.MessageCodec {
return m
}
+func (m *MessageCodec) classifyTransportError(err error)
transports.TransportErrorKind {
+ if err == nil {
+ return transports.TransportErrorUnknown
+ }
+ if transport := m.GetTransportInstance(); transport != nil {
+ return transport.ClassifyError(err)
+ }
+ return transports.TransportErrorUnknown
+}
+
+func (m *MessageCodec) isFatalTransportError(err error) bool {
+ if err == nil || transports.ErrorIs(err, io.EOF) {
+ return false
+ }
+ return m.classifyTransportError(err) == transports.TransportErrorFatal
+}
+
+func (m *MessageCodec) wrapFatalTransportError(err error, msg string) error {
+ if err == nil {
+ return nil
+ }
+ // TODO: Any additional context?
+ return transports.NewTransportError(transports.TransportErrorFatal,
errors.Wrap(err, msg))
+}
+
func (m *MessageCodec) Send(ctx context.Context, interactionInfo string,
message spi.Message) error {
m.log.Trace().Str("interactionInfo", interactionInfo).Msg("Sending
message")
// Cast the message to the correct type of struct
@@ -96,11 +122,15 @@ func (m *MessageCodec) Receive(ctx context.Context)
(spi.Message, error) {
if err := transportInstance.FillBuffer(ctx, func(pos uint, currentByte
byte, reader transports.ExtendedReader) bool {
numBytesAvailable, err :=
transportInstance.GetNumBytesAvailableInBuffer()
if err != nil {
+ m.log.Warn().Err(err).Msg("error getting bytes while
filling buffer")
return false
}
return numBytesAvailable < 6
}); err != nil {
m.log.Warn().Err(err).Msg("error filling buffer")
+ if m.isFatalTransportError(err) {
+ return nil, m.wrapFatalTransportError(err, "error
filling buffer")
+ }
}
// We need at least 6 bytes in order to know how big the packet is in
total
@@ -109,7 +139,9 @@ func (m *MessageCodec) Receive(ctx context.Context)
(spi.Message, error) {
data, err := transportInstance.PeekReadableBytes(ctx, 6)
if err != nil {
m.log.Warn().Err(err).Msg("error peeking")
- // TODO: Possibly clean up ...
+ if m.isFatalTransportError(err) {
+ return nil, m.wrapFatalTransportError(err,
"error peeking header")
+ }
return nil, nil
}
// Get the size of the entire packet little endian plus size of
header
@@ -123,11 +155,17 @@ func (m *MessageCodec) Receive(ctx context.Context)
(spi.Message, error) {
return numBytesAvailable < packetSize
}); err != nil {
m.log.Warn().Err(err).Msg("error filling
buffer")
+ if m.isFatalTransportError(err) {
+ return nil, errors.Wrap(err, "error
filling buffer")
+ }
}
}
data, err = transportInstance.Read(ctx, packetSize)
if err != nil {
- // TODO: Possibly clean up ...
+ m.log.Warn().Err(err).Msg("error reading packet data")
+ if m.isFatalTransportError(err) {
+ return nil, m.wrapFatalTransportError(err,
"error reading packet data")
+ }
return nil, nil
}
rb := utils.NewReadBufferByteBased(data,
utils.WithByteOrderForReadBufferByteBased(binary.LittleEndian))
@@ -140,6 +178,9 @@ func (m *MessageCodec) Receive(ctx context.Context)
(spi.Message, error) {
return tcpPacket, nil
} else if err != nil {
m.log.Warn().Err(err).Msg("Got error reading")
+ if m.isFatalTransportError(err) {
+ return nil, m.wrapFatalTransportError(err, "error
getting readable bytes")
+ }
return nil, nil
}
// TODO: maybe we return here a not enough error error
diff --git a/plc4go/internal/bacnetip/MessageCodec_test.go
b/plc4go/internal/bacnetip/MessageCodec_test.go
index 4036fce80d..e55bb67d34 100644
--- a/plc4go/internal/bacnetip/MessageCodec_test.go
+++ b/plc4go/internal/bacnetip/MessageCodec_test.go
@@ -122,6 +122,10 @@ func (f *fakeTransportInstance) FillBuffer(ctx
context.Context, until func(pos u
return ctx.Err()
}
+func (f *fakeTransportInstance) ClassifyError(err error)
transports.TransportErrorKind {
+ return transports.TransportErrorUnknown
+}
+
var _ transports.TransportInstance = (*fakeTransportInstance)(nil)
func newTestCodec(t *testing.T) (*MessageCodec, *fakeTransportInstance) {
diff --git a/plc4go/internal/cbus/MessageCodec.go
b/plc4go/internal/cbus/MessageCodec.go
index 6d2091489f..2dcd8fba9b 100644
--- a/plc4go/internal/cbus/MessageCodec.go
+++ b/plc4go/internal/cbus/MessageCodec.go
@@ -22,6 +22,7 @@ package cbus
import (
"context"
"hash/crc32"
+ "reflect"
"sync"
"sync/atomic"
"time"
@@ -77,6 +78,16 @@ func (m *MessageCodec) GetCodec() spi.MessageCodec {
return m
}
+func (m *MessageCodec) SetTransportErrorHandler(handler
transports.TransportErrorHandler) {
+ if m.DefaultCodec == nil {
+ return
+ }
+ if value := reflect.ValueOf(m.DefaultCodec); value.Kind() ==
reflect.Pointer && value.IsNil() {
+ return
+ }
+ m.DefaultCodec.SetTransportErrorHandler(handler)
+}
+
func (m *MessageCodec) Connect(ctx context.Context) error {
m.stateChange.Lock()
defer m.stateChange.Unlock()
diff --git a/plc4go/internal/knxnetip/Connection.go
b/plc4go/internal/knxnetip/Connection.go
index 29f9cff328..4a26b3b491 100644
--- a/plc4go/internal/knxnetip/Connection.go
+++ b/plc4go/internal/knxnetip/Connection.go
@@ -28,6 +28,7 @@ import (
"strconv"
"strings"
"sync"
+ "sync/atomic"
"time"
"github.com/rs/zerolog"
@@ -140,6 +141,8 @@ type Connection struct {
passLogToModel bool
log zerolog.Logger
_options []options.WithOption // Used to pass them downstream
+
+ invalidated atomic.Bool
}
var (
@@ -229,6 +232,9 @@ func (m *Connection) GetTracer() tracer.Tracer {
}
func (m *Connection) Connect(ctx context.Context) error {
+ // Reset invalidation state before we start a new connection attempt.
+ m.invalidated.Store(false)
+
// Open the UDP Connection
err := m.messageCodec.Connect(ctx)
if err != nil {
@@ -364,6 +370,7 @@ func (m *Connection) Connect(ctx context.Context) error {
return m.doSomethingAndClose(func() error { return
errors.New("this device doesn't support tunneling") })
}
+ m.invalidated.Store(false)
return nil
}
@@ -419,6 +426,9 @@ func (m *Connection) IsConnected() bool {
}
func (m *Connection) Ping(ctx context.Context) error {
+ if m.IsInvalidated() {
+ return errors.New("connection has been invalidated")
+ }
// Send the connection state request
if _, err := m.sendConnectionStateRequest(ctx); err != nil {
return errors.Wrap(err, "got an error")
@@ -426,6 +436,20 @@ func (m *Connection) Ping(ctx context.Context) error {
return nil
}
+func (m *Connection) Invalidate() {
+ if m.invalidated.Swap(true) {
+ return
+ }
+ m.log.Debug().Msg("invalidating connection")
+ if err := m.Close(); err != nil {
+ m.log.Warn().Err(err).Msg("error closing invalidated
connection")
+ }
+}
+
+func (m *Connection) IsInvalidated() bool {
+ return m.invalidated.Load()
+}
+
func (m *Connection) GetMetadata() apiModel.PlcConnectionMetadata {
return m.metadata
}
diff --git a/plc4go/internal/modbus/Connection.go
b/plc4go/internal/modbus/Connection.go
index dfaff54a32..ba2f2da117 100644
--- a/plc4go/internal/modbus/Connection.go
+++ b/plc4go/internal/modbus/Connection.go
@@ -26,11 +26,11 @@ import (
"github.com/rs/zerolog"
- "github.com/apache/plc4x/plc4go/pkg/api"
+ plc4go "github.com/apache/plc4x/plc4go/pkg/api"
apiModel "github.com/apache/plc4x/plc4go/pkg/api/model"
readWriteModel
"github.com/apache/plc4x/plc4go/protocols/modbus/readwrite/model"
"github.com/apache/plc4x/plc4go/spi"
- "github.com/apache/plc4x/plc4go/spi/default"
+ _default "github.com/apache/plc4x/plc4go/spi/default"
"github.com/apache/plc4x/plc4go/spi/errors"
"github.com/apache/plc4x/plc4go/spi/interceptors"
spiModel "github.com/apache/plc4x/plc4go/spi/model"
@@ -110,6 +110,9 @@ func (c *Connection) GetMessageCodec() spi.MessageCodec {
}
func (c *Connection) Ping(ctx context.Context) error {
+ if c.DefaultConnection.IsInvalidated() {
+ return errors.New("connection has been invalidated")
+ }
c.log.Trace().Msg("Pinging")
errChan := make(chan error, 1)
successChan := make(chan struct{}, 1)
diff --git a/plc4go/internal/modbus/MessageCodec.go
b/plc4go/internal/modbus/MessageCodec.go
index 1c388e7987..3f39e98c3e 100644
--- a/plc4go/internal/modbus/MessageCodec.go
+++ b/plc4go/internal/modbus/MessageCodec.go
@@ -66,6 +66,38 @@ func (m *MessageCodec) GetCodec() spi.MessageCodec {
return m
}
+func (m *MessageCodec) classifyTransportError(err error)
transports.TransportErrorKind {
+ if err == nil {
+ return transports.TransportErrorUnknown
+ }
+ if transport := m.GetTransportInstance(); transport != nil {
+ return transport.ClassifyError(err)
+ }
+ return transports.TransportErrorUnknown
+}
+
+func (m *MessageCodec) isFatalTransportError(err error) bool {
+ if err == nil || isEOF(err) {
+ return false
+ }
+ return m.classifyTransportError(err) == transports.TransportErrorFatal
+}
+
+func (m *MessageCodec) wrapFatalTransportError(err error, msg string) error {
+ if err == nil {
+ return nil
+ }
+ // TODO: Any additional context?
+ return transports.NewTransportError(transports.TransportErrorFatal,
errors.Wrap(err, msg))
+}
+
+func isEOF(err error) (matched bool) {
+ if err == nil {
+ return false
+ }
+ return transports.ErrorIs(err, io.EOF)
+}
+
func (m *MessageCodec) Send(ctx context.Context, interactionInfo string,
message spi.Message) error {
m.log.Trace().Str("interactionInfo", interactionInfo).Msg("Sending
message")
// Cast the message to the correct type of struct
@@ -100,20 +132,23 @@ func (m *MessageCodec) Receive(ctx context.Context)
(spi.Message, error) {
}
return numBytesAvailable < 6
}); err != nil {
- if err != io.EOF {
+ if m.isFatalTransportError(err) {
m.log.Debug().Err(err).Msg("error filling buffer")
+ return nil, m.wrapFatalTransportError(err, "error
filling buffer")
}
- // Fall through on errors, we might have enough data...
+ // Fall through on non-fatal errors, we might have enough
data...
}
// 2. Check buffer status
numBytesAvail, err := ti.GetNumBytesAvailableInBuffer()
- if err != nil {
+ if err != nil && numBytesAvail < 6 {
// Yield if we can't check buffer
- if err == io.EOF {
+ if isEOF(err) {
+ m.log.Debug().Msg("transport buffer exhausted while
checking availability")
return nil, nil
}
- return nil, fmt.Errorf("error getting buffer length")
+ m.log.Warn().Err(err).Msg("error getting buffer length")
+ return nil, m.wrapFatalTransportError(err, "error getting
buffer length")
}
// Need at least 6 bytes for MBAP header
@@ -125,6 +160,9 @@ func (m *MessageCodec) Receive(ctx context.Context)
(spi.Message, error) {
header, err := ti.PeekReadableBytes(ctx, 6)
if err != nil {
m.log.Warn().Err(err).Msg("error peeking header")
+ if m.isFatalTransportError(err) {
+ return nil, m.wrapFatalTransportError(err, "error
peeking header")
+ }
return nil, nil
}
@@ -175,7 +213,10 @@ func (m *MessageCodec) Receive(ctx context.Context)
(spi.Message, error) {
// Read the entire frame
frameSlice, err := ti.PeekReadableBytes(ctx, packetSize)
if err != nil {
- m.log.Warn().Err(err).Msg("Error peeking frame slice")
+ m.log.Warn().Err(err).Msg("error peeking frame slice")
+ if m.isFatalTransportError(err) {
+ return nil, m.wrapFatalTransportError(err, "error
peeking frame slice")
+ }
return nil, nil
}
@@ -357,7 +398,7 @@ func (m *MessageCodec) checkPacketConsistency(data []byte)
bool {
// handleDesync handles stream realignment when an invalid header is detected
at the head.
// It strictly scans the available buffer for a valid MBAP header using
checkPacketConsistency.
// If one is found, it realigns the stream.
-// If NO valid header is found in the *entire* buffer, it treats the
connection as dead.
+// If NO valid header is found in the *entire* buffer, it trims the garbage
and waits for more data.
func (m *MessageCodec) handleDesync(ctx context.Context, reason string, fields
map[string]interface{}) (spi.Message, error) {
ti := m.GetTransportInstance()
@@ -444,5 +485,5 @@ func (m *MessageCodec) handleDesync(ctx context.Context,
reason string, fields m
return nil, err // Return error to kill connection
}
- return nil, fmt.Errorf("stream desynchronized: discarded %d bytes of
garbage", bytesToDiscard)
+ return nil, nil
}
diff --git a/plc4go/internal/simulated/Connection.go
b/plc4go/internal/simulated/Connection.go
index de2ed09600..b18a9e0c4b 100644
--- a/plc4go/internal/simulated/Connection.go
+++ b/plc4go/internal/simulated/Connection.go
@@ -23,6 +23,7 @@ import (
"context"
"strconv"
"sync"
+ "sync/atomic"
"time"
"github.com/rs/zerolog"
@@ -45,6 +46,7 @@ type Connection struct {
connected bool
connectionId string
tracer tracer.Tracer
+ invalidated atomic.Bool
wg sync.WaitGroup // use to track spawned go routines
@@ -118,6 +120,7 @@ func (c *Connection) Connect(_ context.Context) error {
} else {
// Mark the connection as "connected"
c.connected = true
+ c.invalidated.Store(false)
if c.tracer != nil {
c.tracer.AddTransactionalTrace(txId, "connect",
"success")
}
@@ -132,6 +135,9 @@ func (c *Connection) Close() error {
// Check if the connection is connected.
if !c.connected {
+ if c.invalidated.Load() {
+ return nil
+ }
if c.tracer != nil {
c.tracer.AddTrace("close", "error: not connected")
}
@@ -161,10 +167,13 @@ func (c *Connection) Close() error {
}
func (c *Connection) IsConnected() bool {
- return c.connected
+ return c.connected && !c.IsInvalidated()
}
func (c *Connection) Ping(ctx context.Context) error {
+ if c.IsInvalidated() {
+ return errors.New("connection has been invalidated")
+ }
// Check if the connection is connected
if !c.connected {
if c.tracer != nil {
@@ -246,3 +255,17 @@ func (c *Connection) BrowseRequestBuilder()
apiModel.PlcBrowseRequestBuilder {
func (c *Connection) String() string {
return "simulatedConnection"
}
+
+func (c *Connection) Invalidate() {
+ if c.invalidated.Swap(true) {
+ return
+ }
+ c.log.Debug().Msg("invalidating connection")
+ if err := c.Close(); err != nil {
+ c.log.Warn().Err(err).Msg("error closing invalidated
connection")
+ }
+}
+
+func (c *Connection) IsInvalidated() bool {
+ return c.invalidated.Load()
+}
diff --git a/plc4go/internal/simulated/Connection_test.go
b/plc4go/internal/simulated/Connection_test.go
index 52e79df9a6..f325ecb800 100644
--- a/plc4go/internal/simulated/Connection_test.go
+++ b/plc4go/internal/simulated/Connection_test.go
@@ -26,6 +26,7 @@ import (
"time"
"github.com/stretchr/testify/assert"
+ "github.com/stretchr/testify/require"
apiModel "github.com/apache/plc4x/plc4go/pkg/api/model"
"github.com/apache/plc4x/plc4go/spi"
@@ -40,6 +41,7 @@ func TestConnection_Connect(t *testing.T) {
valueHandler spi.PlcValueHandler
options map[string][]string
connected bool
+ invalidated bool
}
tests := []struct {
name string
@@ -93,6 +95,7 @@ func TestConnection_Close(t *testing.T) {
valueHandler spi.PlcValueHandler
options map[string][]string
connected bool
+ invalidated bool
}
tests := []struct {
name string
@@ -291,6 +294,7 @@ func TestConnection_IsConnected(t *testing.T) {
valueHandler spi.PlcValueHandler
options map[string][]string
connected bool
+ invalidated bool
}
tests := []struct {
name string
@@ -319,6 +323,18 @@ func TestConnection_IsConnected(t *testing.T) {
},
want: false,
},
+ {
+ name: "invalidated",
+ fields: fields{
+ device: NewDevice("hurz"),
+ fieldHandler: NewTagHandler(),
+ valueHandler: NewValueHandler(),
+ options: map[string][]string{},
+ connected: true,
+ invalidated: true,
+ },
+ want: false,
+ },
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
@@ -329,6 +345,9 @@ func TestConnection_IsConnected(t *testing.T) {
options: tt.fields.options,
connected: tt.fields.connected,
}
+ if tt.fields.invalidated {
+ c.invalidated.Store(true)
+ }
if got := c.IsConnected(); got != tt.want {
t.Errorf("IsConnected() = %v, want %v", got,
tt.want)
}
@@ -351,6 +370,7 @@ func TestConnection_Ping(t *testing.T) {
name string
fields fields
args args
+ prepare func(*Connection)
wantErr assert.ErrorAssertionFunc
delayAtLeast time.Duration
}{
@@ -386,6 +406,24 @@ func TestConnection_Ping(t *testing.T) {
wantErr: assert.NoError,
delayAtLeast: 1000,
},
+ {
+ name: "invalidated",
+ fields: fields{
+ device: NewDevice("hurz"),
+ fieldHandler: NewTagHandler(),
+ valueHandler: NewValueHandler(),
+ options: map[string][]string{},
+ connected: true,
+ },
+ args: args{
+ ctx: t.Context(),
+ },
+ prepare: func(c *Connection) {
+ c.invalidated.Store(true)
+ },
+ wantErr: assert.Error,
+ delayAtLeast: 0,
+ },
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
@@ -396,12 +434,28 @@ func TestConnection_Ping(t *testing.T) {
options: tt.fields.options,
connected: tt.fields.connected,
}
+ if tt.prepare != nil {
+ prepare := tt.prepare
+ prepare(c)
+ }
err := c.Ping(tt.args.ctx)
tt.wantErr(t, err)
})
}
}
+func TestConnection_Invalidate(t *testing.T) {
+ conn := NewConnection(NewDevice("hurz"), NewTagHandler(),
NewValueHandler(), map[string][]string{})
+ require.NoError(t, conn.Connect(t.Context()))
+ conn.Invalidate()
+ assert.True(t, conn.IsInvalidated())
+ assert.False(t, conn.IsConnected())
+ assert.Error(t, conn.Ping(t.Context()))
+ require.NoError(t, conn.Close())
+ conn.Invalidate()
+ assert.True(t, conn.IsInvalidated())
+}
+
func TestConnection_BrowseRequestBuilder(t *testing.T) {
type fields struct {
device *Device
diff --git a/plc4go/pkg/api/PlcConnection.go b/plc4go/pkg/api/PlcConnection.go
index c24cdd2697..07f16201d1 100644
--- a/plc4go/pkg/api/PlcConnection.go
+++ b/plc4go/pkg/api/PlcConnection.go
@@ -39,6 +39,9 @@ type PlcConnection interface {
// Ping Executes a no-op operation to check if the current connection
is still able to communicate
Ping(ctx context.Context) error
+ // Invalidate marks the connection as irrecoverably failed so caches
can drop it without health checks.
+ Invalidate()
+
// GetMetadata Get some metadata regarding the current connection
GetMetadata() model.PlcConnectionMetadata
diff --git a/plc4go/pkg/api/cache/plcConnectionLease.go
b/plc4go/pkg/api/cache/plcConnectionLease.go
index cd78b88167..8094ec31b8 100644
--- a/plc4go/pkg/api/cache/plcConnectionLease.go
+++ b/plc4go/pkg/api/cache/plcConnectionLease.go
@@ -22,6 +22,7 @@ package cache
import (
"context"
"fmt"
+ "sync/atomic"
"time"
apiModel "github.com/apache/plc4x/plc4go/pkg/api/model"
@@ -38,8 +39,12 @@ type plcConnectionLease struct {
connection tracedPlcConnection
// the last traces of this connection
lastTraces []tracer.TraceEntry
+ // invalidated indicates the lease was explicitly marked unusable by
the caller.
+ invalidated atomic.Bool
}
+var errConnectionInvalidated = errors.New("connection has been invalidated")
+
func newPlcConnectionLease(connectionContainer *connectionContainer, leaseId
uint32, connection tracedPlcConnection) *plcConnectionLease {
p := &plcConnectionLease{
connectionContainer: connectionContainer,
@@ -53,23 +58,17 @@ func newPlcConnectionLease(connectionContainer
*connectionContainer, leaseId uin
}
func (t *plcConnectionLease) IsTraceEnabled() bool {
- if t.connection == nil {
- panic("Called 'IsTraceEnabled' on a closed cached connection")
- }
+ t.ensureConnection("IsTraceEnabled")
return t.connection.IsTraceEnabled()
}
func (t *plcConnectionLease) GetTracer() tracer.Tracer {
- if t.connection == nil {
- panic("Called 'GetTracer' on a closed cached connection")
- }
+ t.ensureConnection("GetTracer")
return t.connection.GetTracer()
}
func (t *plcConnectionLease) GetConnectionId() string {
- if t.connection == nil {
- panic("Called 'GetConnectionId' on a closed cached connection")
- }
+ t.ensureConnection("GetConnectionId")
return fmt.Sprintf("%s-%d", t.connection.GetConnectionId(), t.leaseId)
}
@@ -88,7 +87,9 @@ func (t *plcConnectionLease) Close() error {
// Check if the connection is still alive, if it is, put it back into
the cache
newState := StateIdle
- if err := t.Ping(ctx); err != nil {
+ if t.isInvalidated() {
+ newState = StateInvalid
+ } else if err := t.Ping(ctx); err != nil {
if errors.Is(err, context.DeadlineExceeded) {
// Add some trace information
if t.connection.IsTraceEnabled() {
@@ -98,8 +99,8 @@ func (t *plcConnectionLease) Close() error {
newState = StateInvalid
}
- // Extract the trace entries from the connection.
- if t.IsTraceEnabled() {
+ // Extract the trace entries from the connection unless it was
invalidated.
+ if !t.isInvalidated() && t.IsTraceEnabled() {
_tracer := t.GetTracer()
// Save all traces.
t.lastTraces = _tracer.GetTraces()
@@ -126,6 +127,9 @@ func (t *plcConnectionLease) IsConnected() bool {
if t.connection == nil {
return false
}
+ if t.isInvalidated() {
+ return false
+ }
return t.connection.IsConnected()
}
@@ -133,49 +137,64 @@ func (t *plcConnectionLease) Ping(ctx context.Context)
error {
if t.connection == nil {
panic("Called 'Ping' on a closed cached connection")
}
+ if t.isInvalidated() {
+ return errConnectionInvalidated
+ }
return t.connection.Ping(ctx)
}
func (t *plcConnectionLease) GetMetadata() apiModel.PlcConnectionMetadata {
- if t.connection == nil {
- panic("Called 'GetMetadata' on a closed cached connection")
- }
+ t.ensureConnection("GetMetadata")
return t.connection.GetMetadata()
}
func (t *plcConnectionLease) ReadRequestBuilder()
apiModel.PlcReadRequestBuilder {
- if t.connection == nil {
- panic("Called 'ReadRequestBuilder' on a closed cached
connection")
- }
+ t.ensureConnection("ReadRequestBuilder")
return t.connection.ReadRequestBuilder()
}
func (t *plcConnectionLease) WriteRequestBuilder()
apiModel.PlcWriteRequestBuilder {
- if t.connection == nil {
- panic("Called 'WriteRequestBuilder' on a closed cached
connection")
- }
+ t.ensureConnection("WriteRequestBuilder")
return t.connection.WriteRequestBuilder()
}
func (t *plcConnectionLease) SubscriptionRequestBuilder()
apiModel.PlcSubscriptionRequestBuilder {
- if t.connection == nil {
- panic("Called 'SubscriptionRequestBuilder' on a closed cached
connection")
- }
+ t.ensureConnection("SubscriptionRequestBuilder")
return t.connection.SubscriptionRequestBuilder()
}
func (t *plcConnectionLease) UnsubscriptionRequestBuilder()
apiModel.PlcUnsubscriptionRequestBuilder {
- if t.connection == nil {
- panic("Called 'UnsubscriptionRequestBuilder' on a closed cached
connection")
- }
+ t.ensureConnection("UnsubscriptionRequestBuilder")
return t.connection.UnsubscriptionRequestBuilder()
}
func (t *plcConnectionLease) BrowseRequestBuilder()
apiModel.PlcBrowseRequestBuilder {
+ t.ensureConnection("BrowseRequestBuilder")
+ return t.connection.BrowseRequestBuilder()
+}
+
+func (t *plcConnectionLease) Invalidate() {
if t.connection == nil {
- panic("Called 'BrowseRequestBuilder' on a closed cached
connection")
+ t.invalidated.Store(true)
+ return
}
- return t.connection.BrowseRequestBuilder()
+ if t.invalidated.Swap(true) {
+ return
+ }
+ t.connection.Invalidate()
+}
+
+func (t *plcConnectionLease) ensureConnection(method string) {
+ if t.connection == nil {
+ panic(fmt.Sprintf("Called '%s' on a closed cached connection",
method))
+ }
+ if t.isInvalidated() {
+ panic(fmt.Sprintf("Called '%s' on an invalidated cached
connection", method))
+ }
+}
+
+func (t *plcConnectionLease) isInvalidated() bool {
+ return t.invalidated.Load()
}
func (t *plcConnectionLease) String() string {
diff --git a/plc4go/pkg/api/cache/plcConnectionLease_test.go
b/plc4go/pkg/api/cache/plcConnectionLease_test.go
index 8f01210133..66343a2a12 100644
--- a/plc4go/pkg/api/cache/plcConnectionLease_test.go
+++ b/plc4go/pkg/api/cache/plcConnectionLease_test.go
@@ -20,17 +20,22 @@
package cache
import (
+ "context"
+ "errors"
"sync"
"testing"
"time"
"github.com/stretchr/testify/assert"
+ "github.com/stretchr/testify/require"
"github.com/apache/plc4x/plc4go/internal/simulated"
plc4go "github.com/apache/plc4x/plc4go/pkg/api"
"github.com/apache/plc4x/plc4go/pkg/api/config"
+ apiModel "github.com/apache/plc4x/plc4go/pkg/api/model"
"github.com/apache/plc4x/plc4go/spi/options"
"github.com/apache/plc4x/plc4go/spi/testutils"
+ "github.com/apache/plc4x/plc4go/spi/tracer"
)
func TestLeasedPlcConnection_IsTraceEnabled(t *testing.T) {
@@ -321,6 +326,49 @@ func TestLeasedPlcConnection_GetMetadata(t *testing.T) {
}
}
+func TestLeasedPlcConnection_InvalidateSkipsPingOnClose(t *testing.T) {
+ logger := testutils.ProduceTestingLogger(t)
+ container := &connectionContainer{
+ lock: &sync.RWMutex{},
+ connectionString: "dummy://invalidate",
+ log: logger,
+ }
+ dummyConn := &dummyTracedConnection{}
+ container.connection = dummyConn
+ container.state = StateIdle
+ container.driverManager = &dummyDriverManager{
+ factory: func() plc4go.PlcConnection { return
&dummyTracedConnection{} },
+ }
+ lease := newPlcConnectionLease(container, 1, dummyConn)
+
+ lease.Invalidate()
+ require.True(t, dummyConn.invalidated)
+
+ require.NoError(t, lease.Close())
+ require.Zero(t, dummyConn.pingCount, "ping should not be invoked when
lease invalidated")
+}
+
+func TestLeasedPlcConnection_PingAfterInvalidate(t *testing.T) {
+ logger := testutils.ProduceTestingLogger(t)
+ container := &connectionContainer{
+ lock: &sync.RWMutex{},
+ connectionString: "dummy://ping",
+ log: logger,
+ }
+ dummyConn := &dummyTracedConnection{}
+ container.connection = dummyConn
+ container.state = StateIdle
+ container.driverManager = &dummyDriverManager{
+ factory: func() plc4go.PlcConnection { return
&dummyTracedConnection{} },
+ }
+ lease := newPlcConnectionLease(container, 1, dummyConn)
+
+ lease.Invalidate()
+
+ err := lease.Ping(context.Background())
+ assert.ErrorIs(t, err, errConnectionInvalidated)
+}
+
func TestLeasedPlcConnection_ReadRequestBuilder(t *testing.T) {
logger := testutils.ProduceTestingLogger(t)
driverManager :=
plc4go.NewPlcDriverManager(config.WithCustomLogger(logger))
@@ -403,6 +451,75 @@ func TestLeasedPlcConnection_WriteRequestBuilder(t
*testing.T) {
}
}
+type dummyDriverManager struct {
+ factory func() plc4go.PlcConnection
+}
+
+func (d *dummyDriverManager) RegisterDriver(plc4go.PlcDriver) {}
+
+func (d *dummyDriverManager) ListDriverNames() []string { return nil }
+
+func (d *dummyDriverManager) GetDriver(string) (plc4go.PlcDriver, error) {
+ return nil, errors.New("not implemented")
+}
+
+func (d *dummyDriverManager) GetConnection(context.Context, string)
(plc4go.PlcConnection, error) {
+ if d.factory != nil {
+ return d.factory(), nil
+ }
+ return &dummyTracedConnection{}, nil
+}
+
+func (d *dummyDriverManager) Discover(context.Context,
func(apiModel.PlcDiscoveryItem), ...plc4go.WithDiscoveryOption) error {
+ return nil
+}
+
+func (d *dummyDriverManager) Close() error { return nil }
+
+type dummyTracedConnection struct {
+ pingCount int
+ invalidated bool
+}
+
+func (d *dummyTracedConnection) String() string { return "dummy" }
+
+func (d *dummyTracedConnection) Close() error { return nil }
+
+func (d *dummyTracedConnection) Connect(context.Context) error { return nil }
+
+func (d *dummyTracedConnection) IsConnected() bool { return !d.invalidated }
+
+func (d *dummyTracedConnection) Ping(context.Context) error {
+ d.pingCount++
+ return nil
+}
+
+func (d *dummyTracedConnection) Invalidate() {
+ d.invalidated = true
+}
+
+func (d *dummyTracedConnection) GetMetadata() apiModel.PlcConnectionMetadata {
return nil }
+
+func (d *dummyTracedConnection) ReadRequestBuilder()
apiModel.PlcReadRequestBuilder { return nil }
+
+func (d *dummyTracedConnection) WriteRequestBuilder()
apiModel.PlcWriteRequestBuilder { return nil }
+
+func (d *dummyTracedConnection) SubscriptionRequestBuilder()
apiModel.PlcSubscriptionRequestBuilder {
+ return nil
+}
+
+func (d *dummyTracedConnection) UnsubscriptionRequestBuilder()
apiModel.PlcUnsubscriptionRequestBuilder {
+ return nil
+}
+
+func (d *dummyTracedConnection) BrowseRequestBuilder()
apiModel.PlcBrowseRequestBuilder { return nil }
+
+func (d *dummyTracedConnection) GetConnectionId() string { return "dummy" }
+
+func (d *dummyTracedConnection) IsTraceEnabled() bool { return false }
+
+func (d *dummyTracedConnection) GetTracer() tracer.Tracer { return nil }
+
func TestLeasedPlcConnection_SubscriptionRequestBuilder(t *testing.T) {
logger := testutils.ProduceTestingLogger(t)
driverManager :=
plc4go.NewPlcDriverManager(config.WithCustomLogger(logger))
diff --git a/plc4go/pkg/api/transports/transports.go
b/plc4go/pkg/api/transports/transports.go
index 5add1287b9..f81b6c7524 100644
--- a/plc4go/pkg/api/transports/transports.go
+++ b/plc4go/pkg/api/transports/transports.go
@@ -20,23 +20,49 @@
package transports
import (
- "github.com/apache/plc4x/plc4go/pkg/api"
+ plc4go "github.com/apache/plc4x/plc4go/pkg/api"
"github.com/apache/plc4x/plc4go/pkg/api/config"
"github.com/apache/plc4x/plc4go/spi"
"github.com/apache/plc4x/plc4go/spi/options/converter"
+ spiTransports "github.com/apache/plc4x/plc4go/spi/transports"
"github.com/apache/plc4x/plc4go/spi/transports/serial"
"github.com/apache/plc4x/plc4go/spi/transports/tcp"
"github.com/apache/plc4x/plc4go/spi/transports/udp"
)
+// TransportError exposes the SPI transport error wrapper on the public API so
callers can perform
+// errors.As checks and inspect transport severity without importing SPI
internals.
+type TransportError = spiTransports.TransportError
+
+// TransportErrorKind mirrors the SPI severity enumeration so API users can
make decisions based on
+// transport error classification.
+type TransportErrorKind = spiTransports.TransportErrorKind
+
+const (
+ // TransportErrorUnknown represents an error the transport could not
classify; treat as fatal when unsure.
+ TransportErrorUnknown TransportErrorKind =
spiTransports.TransportErrorUnknown
+
+ // TransportErrorTransient signals a short-lived transport hiccup that
usually succeeds if re-tried immediately.
+ TransportErrorTransient TransportErrorKind =
spiTransports.TransportErrorTransient
+
+ // TransportErrorRetryable indicates the caller should retry the
operation after resetting or reconnecting the transport.
+ TransportErrorRetryable TransportErrorKind =
spiTransports.TransportErrorRetryable
+
+ // TransportErrorFatal marks the transport as unusable for further
work; callers must tear down and rebuild the connection.
+ TransportErrorFatal TransportErrorKind =
spiTransports.TransportErrorFatal
+)
+
+// RegisterTcpTransport registers the TCP transport implementation with the
supplied driver manager using the provided options.
func RegisterTcpTransport(driverManager plc4go.PlcDriverManager, _options
...config.WithOption) {
driverManager.(spi.TransportAware).RegisterTransport(tcp.NewTransport(converter.WithOptionToInternal(_options...)...))
}
+// RegisterUdpTransport registers the UDP transport implementation with the
supplied driver manager using the provided options.
func RegisterUdpTransport(driverManager plc4go.PlcDriverManager, _options
...config.WithOption) {
driverManager.(spi.TransportAware).RegisterTransport(udp.NewTransport(converter.WithOptionToInternal(_options...)...))
}
+// RegisterSerialTransport registers the serial transport implementation with
the supplied driver manager using the provided options.
func RegisterSerialTransport(driverManager plc4go.PlcDriverManager, _options
...config.WithOption) {
driverManager.(spi.TransportAware).RegisterTransport(serial.NewTransport(converter.WithOptionToInternal(_options...)...))
}
diff --git a/plc4go/spi/TransportErrorHandler.go
b/plc4go/spi/TransportErrorHandler.go
new file mode 100644
index 0000000000..8002edecec
--- /dev/null
+++ b/plc4go/spi/TransportErrorHandler.go
@@ -0,0 +1,9 @@
+package spi
+
+import "github.com/apache/plc4x/plc4go/spi/transports"
+
+// TransportErrorHandlerSetter exposes the ability to receive notifications
when
+// the underlying transport reports an error classification.
+type TransportErrorHandlerSetter interface {
+ SetTransportErrorHandler(handler transports.TransportErrorHandler)
+}
diff --git a/plc4go/spi/default/DefaultCodec.go
b/plc4go/spi/default/DefaultCodec.go
index e56df867c4..729dc3d5a2 100644
--- a/plc4go/spi/default/DefaultCodec.go
+++ b/plc4go/spi/default/DefaultCodec.go
@@ -21,6 +21,7 @@ package _default
import (
"context"
+ stdErrors "errors"
"runtime/debug"
"slices"
"sync"
@@ -50,6 +51,7 @@ type DefaultCodec interface {
utils.Serializable
spi.MessageCodec
spi.TransportInstanceExposer
+ spi.TransportErrorHandlerSetter
}
// NewDefaultCodec is the factory for a DefaultCodec
@@ -100,6 +102,8 @@ type defaultCodec struct {
wg sync.WaitGroup // use to track spawned go routines
log zerolog.Logger
+
+ transportErrorHandler transports.TransportErrorHandler
}
func buildDefaultCodec(defaultCodecRequirements DefaultCodecRequirements,
transportInstance transports.TransportInstance, _options ...options.WithOption)
DefaultCodec {
@@ -144,6 +148,10 @@ func (m *defaultCodec) GetTransportInstance()
transports.TransportInstance {
return m.transportInstance
}
+func (m *defaultCodec) SetTransportErrorHandler(handler
transports.TransportErrorHandler) {
+ m.transportErrorHandler = handler
+}
+
func (m *defaultCodec) GetDefaultIncomingMessageChannel() chan spi.Message {
return m.defaultIncomingMessageChannel
}
@@ -476,7 +484,11 @@ mainLoop:
}
if err != nil {
workerLog.Error().Err(err).Msg("got an error reading
from transport")
- continue mainLoop
+ if m.handleTransportError(workerLog, err) {
+ continue mainLoop
+ }
+ workerLog.Debug().Msg("transport error requested worker
shutdown")
+ return
}
if message == nil {
workerLog.Trace().Msg("Not enough data yet")
@@ -515,3 +527,98 @@ func (m *defaultCodec)
passToDefaultIncomingMessageChannel(workerLog zerolog.Log
workerLog.Warn().Interface("message", message).Msg("Message
discarded")
}
}
+
+func (m *defaultCodec) handleTransportError(workerLog zerolog.Logger, err
error) bool {
+ if err == nil {
+ return true
+ }
+ if stdErrors.Is(err, context.Canceled) {
+ workerLog.Debug().Msg("receive aborted due to context
cancellation")
+ return false
+ }
+
+ kind := transports.TransportErrorUnknown
+ if stdErrors.Is(err, context.DeadlineExceeded) {
+ kind = transports.TransportErrorRetryable
+ } else if m.transportInstance != nil {
+ kind = m.transportInstance.ClassifyError(err)
+ }
+ if kind == transports.TransportErrorUnknown {
+ workerLog.Warn().Err(err).Msg("transport error classified as
unknown; treating as fatal")
+ kind = transports.TransportErrorFatal
+ }
+
+ switch kind {
+ case transports.TransportErrorTransient:
+ workerLog.Debug().Err(err).Msg("transient transport error;
keeping worker alive")
+ m.emitTransportError(kind, transports.NewTransportError(kind,
err))
+ return true
+ case transports.TransportErrorRetryable:
+ workerLog.Warn().Err(err).Msg("retryable transport error;
resetting transport instance")
+ if m.transportInstance != nil {
+ defer func() {
+ if recoverErr := recover(); recoverErr != nil {
+ workerLog.Error().Interface("panic",
recoverErr).Msg("panic while resetting transport instance")
+ }
+ }()
+ m.transportInstance.Reset()
+ }
+ m.emitTransportError(kind, transports.NewTransportError(kind,
err))
+ return true
+ case transports.TransportErrorFatal:
+ workerLog.Error().Err(err).Msg("fatal transport error; shutting
down codec")
+ wrappedErr := transports.NewTransportError(kind, err)
+ m.failAllExpectations(wrappedErr)
+ if m.transportInstance != nil {
+ if closeErr := m.transportInstance.Close(); closeErr !=
nil {
+ workerLog.Warn().Err(closeErr).Msg("error
closing transport after fatal condition")
+ }
+ }
+ if m.ctxCancel != nil {
+ m.ctxCancel()
+ }
+ m.running.Store(false)
+ m.emitTransportError(kind, wrappedErr)
+ return false
+ default:
+ workerLog.Error().Err(err).Msg("unexpected transport error
classification; treating as fatal")
+ wrappedErr :=
transports.NewTransportError(transports.TransportErrorFatal, err)
+ m.failAllExpectations(wrappedErr)
+ if m.transportInstance != nil {
+ if closeErr := m.transportInstance.Close(); closeErr !=
nil {
+ workerLog.Warn().Err(closeErr).Msg("error
closing transport after unexpected classification")
+ }
+ }
+ if m.ctxCancel != nil {
+ m.ctxCancel()
+ }
+ m.running.Store(false)
+ m.emitTransportError(transports.TransportErrorFatal, wrappedErr)
+ return false
+ }
+}
+
+func (m *defaultCodec) emitTransportError(kind transports.TransportErrorKind,
err error) {
+ if m.transportErrorHandler != nil {
+ m.transportErrorHandler(kind, err)
+ }
+}
+
+func (m *defaultCodec) failAllExpectations(err error) {
+ m.expectationsChangeMutex.Lock()
+ expectations := slices.Clone(m.expectations)
+ m.expectations = nil
+ m.expectationsChangeMutex.Unlock()
+
+ for _, expectation := range expectations {
+ expectation := expectation
+ expectation.Cancel(err)
+ if handleErr := expectation.GetHandleError(); handleErr != nil {
+ m.wg.Go(func() {
+ if handlerErr := handleErr(err); handlerErr !=
nil {
+
m.log.Error().Err(handlerErr).Msg("error returned by expectation error handler")
+ }
+ })
+ }
+ }
+}
diff --git a/plc4go/spi/default/DefaultConnection.go
b/plc4go/spi/default/DefaultConnection.go
index c88917e422..6a41669146 100644
--- a/plc4go/spi/default/DefaultConnection.go
+++ b/plc4go/spi/default/DefaultConnection.go
@@ -21,12 +21,13 @@ package _default
import (
"context"
+ "reflect"
"sync"
"sync/atomic"
"github.com/rs/zerolog"
- "github.com/apache/plc4x/plc4go/pkg/api"
+ plc4go "github.com/apache/plc4x/plc4go/pkg/api"
apiModel "github.com/apache/plc4x/plc4go/pkg/api/model"
"github.com/apache/plc4x/plc4go/spi"
"github.com/apache/plc4x/plc4go/spi/errors"
@@ -51,6 +52,7 @@ type DefaultConnection interface {
spi.TransportInstanceExposer
spi.HandlerExposer
SetConnected(connected bool)
+ IsInvalidated() bool
}
// NewDefaultConnection is the factory for a DefaultConnection
@@ -87,6 +89,7 @@ type defaultConnection struct {
DefaultConnectionRequirements `ignore:"true"`
// connected indicates if a connection is connected
connected atomic.Bool
+ invalidated atomic.Bool
tagHandler spi.PlcTagHandler
valueHandler spi.PlcValueHandler
@@ -109,18 +112,50 @@ func buildDefaultConnection(requirements
DefaultConnectionRequirements, _options
}
customLogger :=
options.ExtractCustomLoggerOrDefaultToGlobal(_options...)
- return &defaultConnection{
+ conn := &defaultConnection{
DefaultConnectionRequirements: requirements,
tagHandler: tagHandler,
valueHandler: valueHandler,
log: customLogger,
}
+
+ var codec spi.MessageCodec
+ if requirements != nil {
+ codec = requirements.GetMessageCodec()
+ }
+ if codec != nil {
+ if codecValue := reflect.ValueOf(codec); codecValue.Kind() ==
reflect.Pointer && codecValue.IsNil() {
+ codec = nil
+ }
+ }
+ if codec != nil {
+ if setter, ok := codec.(spi.TransportErrorHandlerSetter); ok {
+ setter.SetTransportErrorHandler(func(kind
transports.TransportErrorKind, err error) {
+ switch kind {
+ case transports.TransportErrorFatal:
+
conn.log.Error().Err(err).Msg("transport reported fatal error; invalidating
connection")
+ conn.Invalidate()
+ case transports.TransportErrorRetryable:
+ conn.log.Warn().Err(err).Msg("transport
reported retryable error")
+ case transports.TransportErrorTransient:
+
conn.log.Debug().Err(err).Msg("transport reported transient error")
+ default:
+ conn.log.Warn().Err(err).Msg("transport
reported unknown error classification")
+ }
+ })
+ }
+ }
+
+ return conn
}
func (d *defaultConnection) SetConnected(connected bool) {
d.log.Trace().Bool("connected", connected).Msg("set connected")
d.connected.Store(connected)
+ if connected {
+ d.invalidated.Store(false)
+ }
}
func (d *defaultConnection) Connect(ctx context.Context) error {
@@ -135,7 +170,11 @@ func (d *defaultConnection) Close() error {
if messageCodec := d.GetMessageCodec(); messageCodec != nil {
d.log.Trace().Msg("disconnecting message codec")
if err := messageCodec.Disconnect(); err != nil {
- d.log.Warn().Err(err).Msg("Error disconnecting message
code")
+ if err.Error() != "already disconnected" {
+ d.log.Warn().Err(err).Msg("Error disconnecting
message codec")
+ } else {
+ d.log.Trace().Msg("message codec already
disconnected")
+ }
} else {
d.log.Trace().Msg("message codec disconnected")
}
@@ -155,16 +194,33 @@ func (d *defaultConnection) Close() error {
func (d *defaultConnection) IsConnected() bool {
// TODO: should we check here if the transport is connected?
- return d.connected.Load()
+ return d.connected.Load() && !d.invalidated.Load()
}
func (d *defaultConnection) Ping(_ context.Context) error {
+ if d.invalidated.Load() {
+ return errors.New("connection has been invalidated")
+ }
if !d.DefaultConnectionRequirements.IsConnected() {
return errors.New("not connected")
}
return nil
}
+func (d *defaultConnection) Invalidate() {
+ if d.invalidated.Swap(true) {
+ return
+ }
+ d.log.Debug().Msg("invalidating connection")
+ if err := d.Close(); err != nil {
+ d.log.Warn().Err(err).Msg("error closing invalidated
connection")
+ }
+}
+
+func (d *defaultConnection) IsInvalidated() bool {
+ return d.invalidated.Load()
+}
+
func (d *defaultConnection) GetMetadata() apiModel.PlcConnectionMetadata {
return &DefaultConnectionMetadata{
ConnectionAttributes: nil,
diff --git a/plc4go/spi/default/mocks_test.go b/plc4go/spi/default/mocks_test.go
index 9b37620105..08ac97bd52 100644
--- a/plc4go/spi/default/mocks_test.go
+++ b/plc4go/spi/default/mocks_test.go
@@ -771,6 +771,46 @@ func (_c *MockDefaultCodec_GetTransportInstance_Call)
RunAndReturn(run func() tr
return _c
}
+// SetTransportErrorHandler provides a mock function for the type
MockDefaultCodec
+func (_mock *MockDefaultCodec) SetTransportErrorHandler(handler
transports.TransportErrorHandler) {
+ _mock.Called(handler)
+ return
+}
+
+// MockDefaultCodec_SetTransportErrorHandler_Call is a *mock.Call that shadows
Run/Return methods with type explicit version for method
'SetTransportErrorHandler'
+type MockDefaultCodec_SetTransportErrorHandler_Call struct {
+ *mock.Call
+}
+
+// SetTransportErrorHandler is a helper method to define mock.On call
+// - handler transports.TransportErrorHandler
+func (_e *MockDefaultCodec_Expecter) SetTransportErrorHandler(handler
interface{}) *MockDefaultCodec_SetTransportErrorHandler_Call {
+ return &MockDefaultCodec_SetTransportErrorHandler_Call{Call:
_e.mock.On("SetTransportErrorHandler", handler)}
+}
+
+func (_c *MockDefaultCodec_SetTransportErrorHandler_Call) Run(run func(handler
transports.TransportErrorHandler))
*MockDefaultCodec_SetTransportErrorHandler_Call {
+ _c.Call.Run(func(args mock.Arguments) {
+ var arg0 transports.TransportErrorHandler
+ if args[0] != nil {
+ arg0 = args[0].(transports.TransportErrorHandler)
+ }
+ run(
+ arg0,
+ )
+ })
+ return _c
+}
+
+func (_c *MockDefaultCodec_SetTransportErrorHandler_Call) Return()
*MockDefaultCodec_SetTransportErrorHandler_Call {
+ _c.Call.Return()
+ return _c
+}
+
+func (_c *MockDefaultCodec_SetTransportErrorHandler_Call) RunAndReturn(run
func(handler transports.TransportErrorHandler))
*MockDefaultCodec_SetTransportErrorHandler_Call {
+ _c.Run(run)
+ return _c
+}
+
// IsRunning provides a mock function for the type MockDefaultCodec
func (_mock *MockDefaultCodec) IsRunning() bool {
ret := _mock.Called()
@@ -4527,6 +4567,57 @@ func (_c
*MockTransportInstance_GetNumBytesAvailableInBuffer_Call) RunAndReturn(
return _c
}
+// ClassifyError provides a mock function for the type MockTransportInstance
+func (_mock *MockTransportInstance) ClassifyError(err error)
transports.TransportErrorKind {
+ ret := _mock.Called(err)
+
+ if len(ret) == 0 {
+ panic("no return value specified for ClassifyError")
+ }
+
+ var r0 transports.TransportErrorKind
+ if returnFunc, ok := ret.Get(0).(func(error)
transports.TransportErrorKind); ok {
+ r0 = returnFunc(err)
+ } else {
+ r0 = ret.Get(0).(transports.TransportErrorKind)
+ }
+ return r0
+}
+
+// MockTransportInstance_ClassifyError_Call is a *mock.Call that shadows
Run/Return methods with type explicit version for method 'ClassifyError'
+type MockTransportInstance_ClassifyError_Call struct {
+ *mock.Call
+}
+
+// ClassifyError is a helper method to define mock.On call
+// - err error
+func (_e *MockTransportInstance_Expecter) ClassifyError(err interface{})
*MockTransportInstance_ClassifyError_Call {
+ return &MockTransportInstance_ClassifyError_Call{Call:
_e.mock.On("ClassifyError", err)}
+}
+
+func (_c *MockTransportInstance_ClassifyError_Call) Run(run func(err error))
*MockTransportInstance_ClassifyError_Call {
+ _c.Call.Run(func(args mock.Arguments) {
+ var arg0 error
+ if args[0] != nil {
+ arg0 = args[0].(error)
+ }
+ run(
+ arg0,
+ )
+ })
+ return _c
+}
+
+func (_c *MockTransportInstance_ClassifyError_Call) Return(kind
transports.TransportErrorKind) *MockTransportInstance_ClassifyError_Call {
+ _c.Call.Return(kind)
+ return _c
+}
+
+func (_c *MockTransportInstance_ClassifyError_Call) RunAndReturn(run func(err
error) transports.TransportErrorKind) *MockTransportInstance_ClassifyError_Call
{
+ _c.Call.Return(run)
+ return _c
+}
+
// IsConnected provides a mock function for the type MockTransportInstance
func (_mock *MockTransportInstance) IsConnected() bool {
ret := _mock.Called()
@@ -4707,6 +4798,46 @@ func (_c *MockTransportInstance_Read_Call)
RunAndReturn(run func(ctx context.Con
return _c
}
+// SetTransportErrorHandler provides a mock function for the type
MockTransportInstance
+func (_mock *MockTransportInstance) SetTransportErrorHandler(handler
transports.TransportErrorHandler) {
+ _mock.Called(handler)
+ return
+}
+
+// MockTransportInstance_SetTransportErrorHandler_Call is a *mock.Call that
shadows Run/Return methods with type explicit version for method
'SetTransportErrorHandler'
+type MockTransportInstance_SetTransportErrorHandler_Call struct {
+ *mock.Call
+}
+
+// SetTransportErrorHandler is a helper method to define mock.On call
+// - handler transports.TransportErrorHandler
+func (_e *MockTransportInstance_Expecter) SetTransportErrorHandler(handler
interface{}) *MockTransportInstance_SetTransportErrorHandler_Call {
+ return &MockTransportInstance_SetTransportErrorHandler_Call{Call:
_e.mock.On("SetTransportErrorHandler", handler)}
+}
+
+func (_c *MockTransportInstance_SetTransportErrorHandler_Call) Run(run
func(handler transports.TransportErrorHandler))
*MockTransportInstance_SetTransportErrorHandler_Call {
+ _c.Call.Run(func(args mock.Arguments) {
+ var arg0 transports.TransportErrorHandler
+ if args[0] != nil {
+ arg0 = args[0].(transports.TransportErrorHandler)
+ }
+ run(
+ arg0,
+ )
+ })
+ return _c
+}
+
+func (_c *MockTransportInstance_SetTransportErrorHandler_Call) Return()
*MockTransportInstance_SetTransportErrorHandler_Call {
+ _c.Call.Return()
+ return _c
+}
+
+func (_c *MockTransportInstance_SetTransportErrorHandler_Call)
RunAndReturn(run func(handler transports.TransportErrorHandler))
*MockTransportInstance_SetTransportErrorHandler_Call {
+ _c.Call.Return(run)
+ return _c
+}
+
// Reset provides a mock function for the type MockTransportInstance
func (_mock *MockTransportInstance) Reset() {
_mock.Called()
diff --git a/plc4go/spi/transports/TransportInstance.go
b/plc4go/spi/transports/TransportInstance.go
index 291523f53a..80535b7a22 100644
--- a/plc4go/spi/transports/TransportInstance.go
+++ b/plc4go/spi/transports/TransportInstance.go
@@ -47,4 +47,7 @@ type TransportInstance interface {
Write(ctx context.Context, data []byte) error
// Reset resets the transport instance
Reset()
+
+ // ClassifyError maps a transport-specific error to a severity to
support recovery decisions.
+ ClassifyError(err error) TransportErrorKind
}
diff --git a/plc4go/spi/transports/errors.go b/plc4go/spi/transports/errors.go
new file mode 100644
index 0000000000..a83f313298
--- /dev/null
+++ b/plc4go/spi/transports/errors.go
@@ -0,0 +1,153 @@
+package transports
+
+import (
+ stdErrors "errors"
+ "fmt"
+ "syscall"
+
+ "github.com/apache/plc4x/plc4go/spi/utils"
+)
+
+// TransportErrorKind represents the transport-level severity of an error.
+type TransportErrorKind int
+
+const (
+ TransportErrorUnknown TransportErrorKind = iota
+ // TransportErrorTransient indicates a transient issue (e.g. temporary
IO hiccup) that may succeed on retry without reconnect.
+ TransportErrorTransient
+ // TransportErrorRetryable indicates the caller should retry the
operation after resetting or reconnecting.
+ TransportErrorRetryable
+ // TransportErrorFatal indicates the transport connection can no longer
be used.
+ TransportErrorFatal
+)
+
+// TransportErrorHandler is invoked when callers classify a transport error.
+type TransportErrorHandler func(kind TransportErrorKind, err error)
+
+// TransportError wraps an underlying error with its classified transport
severity.
+type TransportError struct {
+ kind TransportErrorKind
+ err error
+}
+
+// NewTransportError creates a new TransportError with the given kind and
cause.
+func NewTransportError(kind TransportErrorKind, err error) error {
+ if utils.IsNil(err) {
+ return nil
+ }
+ var existing *TransportError
+ if ErrorAs(err, &existing) {
+ return err
+ }
+ return &TransportError{kind: kind, err: err}
+}
+
+// AsTransportError retrieves a TransportError from the provided error chain.
+func AsTransportError(err error) (*TransportError, bool) {
+ if utils.IsNil(err) {
+ return nil, false
+ }
+ var transportErr *TransportError
+ if ErrorAs(err, &transportErr) {
+ return transportErr, true
+ }
+ return nil, false
+}
+
+// Error implements the error interface.
+func (t *TransportError) Error() string {
+ if t.err == nil {
+ return fmt.Sprintf("transport error (%s)", t.kind.String())
+ }
+ return fmt.Sprintf("transport error (%s): %v", t.kind.String(), t.err)
+}
+
+// Unwrap exposes the underlying cause.
+func (t *TransportError) Unwrap() error {
+ return t.err
+}
+
+// Kind reports the TransportErrorKind associated with the error.
+func (t *TransportError) Kind() TransportErrorKind {
+ return t.kind
+}
+
+// ErrorIs mirrors errors.Is but guards against panics triggered by improperly
constructed error values.
+func ErrorIs(err error, target error) (matched bool) {
+ if err == nil || target == nil || utils.IsNil(err) {
+ return false
+ }
+ defer func() {
+ if r := recover(); r != nil {
+ matched = false
+ }
+ }()
+ return stdErrors.Is(err, target)
+}
+
+// ErrorAs mirrors errors.As but guards against panics triggered by improperly
constructed error values.
+func ErrorAs(err error, target any) (matched bool) {
+ if err == nil || target == nil || utils.IsNil(err) {
+ return false
+ }
+ defer func() {
+ if r := recover(); r != nil {
+ matched = false
+ }
+ }()
+ return stdErrors.As(err, target)
+}
+
+// IsFatal reports whether the error kind signals an unusable transport.
+func (k TransportErrorKind) IsFatal() bool {
+ return k == TransportErrorFatal
+}
+
+// IsRetryable reports whether the operation may succeed if repeated.
+func (k TransportErrorKind) IsRetryable() bool {
+ return k == TransportErrorTransient || k == TransportErrorRetryable
+}
+
+// IsTransientSyscallError checks for errno values that are commonly treated
as transient.
+func IsTransientSyscallError(err error) bool {
+ if err == nil || utils.IsNil(err) {
+ return false
+ }
+ var errno syscall.Errno
+ if !ErrorAs(err, &errno) {
+ return false
+ }
+ switch errno {
+ case syscall.EAGAIN,
+ syscall.EINTR,
+ syscall.EINPROGRESS,
+ syscall.EALREADY,
+ syscall.ENOBUFS:
+ return true
+ }
+ switch errno {
+ case syscall.Errno(10004), // WSAEINTR
+ syscall.Errno(10035), // WSAEWOULDBLOCK
+ syscall.Errno(10036), // WSAEINPROGRESS
+ syscall.Errno(10037), // WSAEALREADY
+ syscall.Errno(10055): // WSAENOBUFS
+ return true
+ }
+ return false
+}
+
+// String returns a human readable representation of the error kind.
+func (k TransportErrorKind) String() string {
+ switch k {
+ case TransportErrorUnknown:
+ return "unknown"
+ case TransportErrorTransient:
+ return "transient"
+ case TransportErrorRetryable:
+ return "retryable"
+ case TransportErrorFatal:
+ return "fatal"
+ default:
+ return "invalid"
+ }
+}
diff --git a/plc4go/spi/transports/mocks_test.go
b/plc4go/spi/transports/mocks_test.go
index 53a6cde7cd..eebffe96c4 100644
--- a/plc4go/spi/transports/mocks_test.go
+++ b/plc4go/spi/transports/mocks_test.go
@@ -751,6 +751,57 @@ func (_c
*MockTransportInstance_GetNumBytesAvailableInBuffer_Call) RunAndReturn(
return _c
}
+// ClassifyError provides a mock function for the type MockTransportInstance
+func (_mock *MockTransportInstance) ClassifyError(err error)
TransportErrorKind {
+ ret := _mock.Called(err)
+
+ if len(ret) == 0 {
+ panic("no return value specified for ClassifyError")
+ }
+
+ var r0 TransportErrorKind
+ if returnFunc, ok := ret.Get(0).(func(error) TransportErrorKind); ok {
+ r0 = returnFunc(err)
+ } else {
+ r0 = ret.Get(0).(TransportErrorKind)
+ }
+ return r0
+}
+
+// MockTransportInstance_ClassifyError_Call is a *mock.Call that shadows
Run/Return methods with type explicit version for method 'ClassifyError'
+type MockTransportInstance_ClassifyError_Call struct {
+ *mock.Call
+}
+
+// ClassifyError is a helper method to define mock.On call
+// - err error
+func (_e *MockTransportInstance_Expecter) ClassifyError(err interface{})
*MockTransportInstance_ClassifyError_Call {
+ return &MockTransportInstance_ClassifyError_Call{Call:
_e.mock.On("ClassifyError", err)}
+}
+
+func (_c *MockTransportInstance_ClassifyError_Call) Run(run func(err error))
*MockTransportInstance_ClassifyError_Call {
+ _c.Call.Run(func(args mock.Arguments) {
+ var arg0 error
+ if args[0] != nil {
+ arg0 = args[0].(error)
+ }
+ run(
+ arg0,
+ )
+ })
+ return _c
+}
+
+func (_c *MockTransportInstance_ClassifyError_Call) Return(kind
TransportErrorKind) *MockTransportInstance_ClassifyError_Call {
+ _c.Call.Return(kind)
+ return _c
+}
+
+func (_c *MockTransportInstance_ClassifyError_Call) RunAndReturn(run func(err
error) TransportErrorKind) *MockTransportInstance_ClassifyError_Call {
+ _c.Call.Return(run)
+ return _c
+}
+
// IsConnected provides a mock function for the type MockTransportInstance
func (_mock *MockTransportInstance) IsConnected() bool {
ret := _mock.Called()
@@ -931,6 +982,46 @@ func (_c *MockTransportInstance_Read_Call)
RunAndReturn(run func(ctx context.Con
return _c
}
+// SetTransportErrorHandler provides a mock function for the type
MockTransportInstance
+func (_mock *MockTransportInstance) SetTransportErrorHandler(handler
TransportErrorHandler) {
+ _mock.Called(handler)
+ return
+}
+
+// MockTransportInstance_SetTransportErrorHandler_Call is a *mock.Call that
shadows Run/Return methods with type explicit version for method
'SetTransportErrorHandler'
+type MockTransportInstance_SetTransportErrorHandler_Call struct {
+ *mock.Call
+}
+
+// SetTransportErrorHandler is a helper method to define mock.On call
+// - handler TransportErrorHandler
+func (_e *MockTransportInstance_Expecter) SetTransportErrorHandler(handler
interface{}) *MockTransportInstance_SetTransportErrorHandler_Call {
+ return &MockTransportInstance_SetTransportErrorHandler_Call{Call:
_e.mock.On("SetTransportErrorHandler", handler)}
+}
+
+func (_c *MockTransportInstance_SetTransportErrorHandler_Call) Run(run
func(handler TransportErrorHandler))
*MockTransportInstance_SetTransportErrorHandler_Call {
+ _c.Call.Run(func(args mock.Arguments) {
+ var arg0 TransportErrorHandler
+ if args[0] != nil {
+ arg0 = args[0].(TransportErrorHandler)
+ }
+ run(
+ arg0,
+ )
+ })
+ return _c
+}
+
+func (_c *MockTransportInstance_SetTransportErrorHandler_Call) Return()
*MockTransportInstance_SetTransportErrorHandler_Call {
+ _c.Call.Return()
+ return _c
+}
+
+func (_c *MockTransportInstance_SetTransportErrorHandler_Call)
RunAndReturn(run func(handler TransportErrorHandler))
*MockTransportInstance_SetTransportErrorHandler_Call {
+ _c.Call.Return(run)
+ return _c
+}
+
// Reset provides a mock function for the type MockTransportInstance
func (_mock *MockTransportInstance) Reset() {
_mock.Called()
diff --git a/plc4go/spi/transports/pcap/TransportInstance.go
b/plc4go/spi/transports/pcap/TransportInstance.go
index 83a98a6f79..8b6790b227 100644
--- a/plc4go/spi/transports/pcap/TransportInstance.go
+++ b/plc4go/spi/transports/pcap/TransportInstance.go
@@ -208,3 +208,16 @@ func (m *TransportInstance) SetReadDeadline(deadline
time.Time) error {
func (m *TransportInstance) String() string {
return fmt.Sprintf("pcap:%s(%s)x%f", m.transportFile, m.portRange,
m.speedFactor)
}
+
+func (m *TransportInstance) ClassifyError(err error)
transports.TransportErrorKind {
+ if err == nil {
+ return transports.TransportErrorUnknown
+ }
+ if transports.ErrorIs(err, io.EOF) {
+ return transports.TransportErrorFatal
+ }
+ if transports.ErrorIs(err, context.Canceled) || transports.ErrorIs(err,
context.DeadlineExceeded) {
+ return transports.TransportErrorTransient
+ }
+ return transports.TransportErrorFatal
+}
diff --git a/plc4go/spi/transports/serial/TransportInstance.go
b/plc4go/spi/transports/serial/TransportInstance.go
index f9c00fae4d..be8cfa0070 100644
--- a/plc4go/spi/transports/serial/TransportInstance.go
+++ b/plc4go/spi/transports/serial/TransportInstance.go
@@ -24,6 +24,7 @@ import (
"context"
"fmt"
"io"
+ "strings"
"sync"
"sync/atomic"
"time"
@@ -154,3 +155,25 @@ func (m *TransportInstance) SetReadDeadline(deadline
time.Time) error {
func (m *TransportInstance) String() string {
return fmt.Sprintf("serial:%s:%d", m.SerialPortName, m.BaudRate)
}
+
+func (m *TransportInstance) ClassifyError(err error)
transports.TransportErrorKind {
+ if err == nil {
+ return transports.TransportErrorUnknown
+ }
+ if transports.IsTransientSyscallError(err) {
+ return transports.TransportErrorTransient
+ }
+ if transports.ErrorIs(err, io.EOF) {
+ return transports.TransportErrorFatal
+ }
+ lower := strings.ToLower(err.Error())
+ switch {
+ case strings.Contains(lower, "timeout"):
+ return transports.TransportErrorRetryable
+ case strings.Contains(lower, "temporarily unavailable"):
+ return transports.TransportErrorTransient
+ case strings.Contains(lower, "closed"):
+ return transports.TransportErrorFatal
+ }
+ return transports.TransportErrorFatal
+}
diff --git a/plc4go/spi/transports/tcp/TransportInstance.go
b/plc4go/spi/transports/tcp/TransportInstance.go
index 5479c49550..ebd014f213 100644
--- a/plc4go/spi/transports/tcp/TransportInstance.go
+++ b/plc4go/spi/transports/tcp/TransportInstance.go
@@ -23,9 +23,11 @@ import (
"bufio"
"context"
"fmt"
+ "io"
"net"
"sync"
"sync/atomic"
+ "syscall"
"time"
"github.com/rs/zerolog"
@@ -158,3 +160,41 @@ func (m *TransportInstance) String() string {
}
return fmt.Sprintf("tcp:%s%s", localAddress, m.RemoteAddress)
}
+
+// ClassifyError attempts to map common network errors to transport severity
categories.
+func (m *TransportInstance) ClassifyError(err error)
transports.TransportErrorKind {
+ if err == nil {
+ return transports.TransportErrorUnknown
+ }
+ if transports.ErrorIs(err, io.EOF) || transports.ErrorIs(err,
net.ErrClosed) {
+ return transports.TransportErrorFatal
+ }
+ if netErr, ok := err.(net.Error); ok {
+ if netErr.Timeout() {
+ return transports.TransportErrorRetryable
+ }
+ }
+ if transports.IsTransientSyscallError(err) {
+ return transports.TransportErrorTransient
+ }
+ var opErr *net.OpError
+ if transports.ErrorAs(err, &opErr) && opErr != nil {
+ if opErr.Timeout() {
+ return transports.TransportErrorRetryable
+ }
+ if transports.IsTransientSyscallError(opErr.Err) {
+ return transports.TransportErrorTransient
+ }
+ var syscallErr syscall.Errno
+ if transports.ErrorAs(opErr.Err, &syscallErr) {
+ switch syscallErr {
+ case syscall.ECONNREFUSED, syscall.ECONNRESET,
syscall.EPIPE, syscall.ENETDOWN, syscall.ENETUNREACH:
+ return transports.TransportErrorFatal
+ case syscall.ETIMEDOUT:
+ return transports.TransportErrorRetryable
+ }
+ }
+ return transports.TransportErrorFatal
+ }
+ return transports.TransportErrorFatal
+}
diff --git a/plc4go/spi/transports/test/TransportInstance.go
b/plc4go/spi/transports/test/TransportInstance.go
index 8965df8c9b..a9e908b9b1 100644
--- a/plc4go/spi/transports/test/TransportInstance.go
+++ b/plc4go/spi/transports/test/TransportInstance.go
@@ -326,6 +326,20 @@ func (m *TransportInstance) String() string {
return "test"
}
+// ClassifyError maps test-transport specific error values to the shared
severity enum.
+func (m *TransportInstance) ClassifyError(err error)
transports.TransportErrorKind {
+ if err == nil {
+ return transports.TransportErrorUnknown
+ }
+ if transports.ErrorIs(err, context.Canceled) {
+ return transports.TransportErrorTransient
+ }
+ if transports.ErrorIs(err, context.DeadlineExceeded) ||
transports.ErrorIs(err, bufio.ErrBufferFull) {
+ return transports.TransportErrorRetryable
+ }
+ return transports.TransportErrorFatal
+}
+
func (m *TransportInstance) availableBytes() uint32 {
m.dataMutex.RLock()
defer m.dataMutex.RUnlock()
diff --git a/plc4go/spi/transports/udp/TransportInstance.go
b/plc4go/spi/transports/udp/TransportInstance.go
index 51703ffc83..e29343a71d 100644
--- a/plc4go/spi/transports/udp/TransportInstance.go
+++ b/plc4go/spi/transports/udp/TransportInstance.go
@@ -23,9 +23,11 @@ import (
"bufio"
"context"
"fmt"
+ "io"
"net"
"sync"
"sync/atomic"
+ "syscall"
"time"
"github.com/rs/zerolog"
@@ -249,3 +251,38 @@ func (m *TransportInstance) Write(ctx context.Context,
data []byte) error {
func (m *TransportInstance) String() string {
return fmt.Sprintf("udp:%s->%s", m.LocalAddress, m.RemoteAddress)
}
+
+func (m *TransportInstance) ClassifyError(err error)
transports.TransportErrorKind {
+ if err == nil {
+ return transports.TransportErrorUnknown
+ }
+ if transports.ErrorIs(err, io.EOF) || transports.ErrorIs(err,
net.ErrClosed) || transports.ErrorIs(err, syscall.EPIPE) {
+ return transports.TransportErrorFatal
+ }
+ if netErr, ok := err.(net.Error); ok {
+ if netErr.Timeout() {
+ return transports.TransportErrorRetryable
+ }
+ }
+ if transports.IsTransientSyscallError(err) {
+ return transports.TransportErrorTransient
+ }
+ var opErr *net.OpError
+ if transports.ErrorAs(err, &opErr) && opErr != nil {
+ if opErr.Timeout() {
+ return transports.TransportErrorRetryable
+ }
+ if transports.IsTransientSyscallError(opErr.Err) {
+ return transports.TransportErrorTransient
+ }
+ if syscallErr, ok := opErr.Err.(syscall.Errno); ok {
+ switch syscallErr {
+ case syscall.ECONNRESET, syscall.ECONNREFUSED,
syscall.ENETDOWN, syscall.ENETUNREACH:
+ return transports.TransportErrorFatal
+ case syscall.ETIMEDOUT:
+ return transports.TransportErrorRetryable
+ }
+ }
+ }
+ return transports.TransportErrorFatal
+}