From 2684093f50d18e303a2de8ffb3f57f5be4df85e5 Mon Sep 17 00:00:00 2001 From: Thomas Maurer Date: Mon, 17 Aug 2026 22:33:28 +0200 Subject: [PATCH] feat: forward vendor function codes mbproxy rejected any function code outside the eight standard ones with Illegal Function. Huawei SUN2000 installer login and optimizer file transfer use 0x41, so wlcrs/huawei_solar could not complete setup through the proxy. Unknown function codes are now forwarded as opaque PDUs. They are not cached and not retried. Read-only mode still applies only to the standard write codes. Fixes #13 Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: 0649ac6f-b014-4168-82f9-d3e4eeeafdfc --- README.md | 6 + SPEC.md | 11 ++ internal/modbus/client.go | 52 ++++++- internal/modbus/client_test.go | 274 +++++++++++++++++++++++++++++++++ internal/modbus/server.go | 5 +- internal/modbus/server_test.go | 69 +++++++++ internal/proxy/proxy.go | 10 +- internal/proxy/proxy_test.go | 119 ++++++++++++-- 8 files changed, 533 insertions(+), 13 deletions(-) diff --git a/README.md b/README.md index aef28fd..54bad38 100644 --- a/README.md +++ b/README.md @@ -7,6 +7,7 @@ A lightweight Modbus TCP proxy with in-memory caching. Designed to reduce load o - **Caching**: In-memory cache with configurable TTL - **Request coalescing**: Identical concurrent requests share a single upstream fetch - **Read-only mode**: Optionally block or ignore write requests +- **Vendor function codes**: Forward non-standard PDUs such as Huawei `0x41` without caching or retrying them - **Auto-reconnect**: Automatic upstream reconnection on failure - **Stale data fallback**: Optionally serve stale cache on upstream errors - **Request diagnostics**: Structured lifecycle timing, retry, exception, and health state @@ -93,6 +94,11 @@ corresponding success has occurred. - `true`: Silently ignore write requests, return success response - `deny`: Reject write requests with Modbus illegal function exception +Read-only mode applies only to the standard write function codes (`0x05`, +`0x06`, `0x0F`, `0x10`). Other function codes, including vendor codes such as +Huawei `0x41`, are forwarded as opaque PDUs. Those requests are not cached and +are not retried. + ## Docker Compose Examples ### Basic Setup diff --git a/SPEC.md b/SPEC.md index 19a88e7..d37a0be 100644 --- a/SPEC.md +++ b/SPEC.md @@ -36,6 +36,9 @@ Many Modbus devices (inverters, meters, battery systems) have limited polling ca - `0x06` Write Single Register - `0x0F` Write Multiple Coils - `0x10` Write Multiple Registers +- Forward any other function code as an opaque PDU. This covers vendor codes + such as Huawei SUN2000 `0x41` (installer login and file transfer). Opaque + requests are not cached and are not retried. ### 2. Upstream Connection - Connect to downstream Modbus device via TCP/IP only @@ -125,6 +128,10 @@ Three modes: - `true` (default): Silently ignore write requests, return success - `deny`: Reject write requests with Modbus exception (illegal function) +Read-only mode applies only to the four standard write function codes. Vendor +and other non-standard function codes are always forwarded, because mbproxy +cannot invent a valid response for an unknown PDU. + ### 5. Graceful Shutdown - Handle SIGTERM/SIGINT signals - Complete in-flight requests before shutdown (with configurable timeout, default: 30s) @@ -275,6 +282,10 @@ The cache also exposes `Coalesce(ctx, rangeKey, fetch)` for request coalescing. - Check readonly mode - If allowed: increment the write generation and invalidate every cached register/coil in the written address range before forwarding upstream - Return response +5. **For other function codes**: + - Forward the raw PDU upstream + - Do not cache, coalesce, or retry + - Do not invent a local success or exception response unless the PDU itself is missing or malformed ## Logging diff --git a/internal/modbus/client.go b/internal/modbus/client.go index 9055e2d..19b2029 100644 --- a/internal/modbus/client.go +++ b/internal/modbus/client.go @@ -29,6 +29,7 @@ type clientSession interface { Close() error SetSlave(byte) BeginRequest(context.Context, time.Duration) func() error + SendRaw(context.Context, []byte) ([]byte, error) } type sessionFactory func() (clientSession, requestClient) @@ -184,6 +185,47 @@ func (c *tcpSession) SetSlave(slaveID byte) { c.handler.SetSlave(slaveID) } +func (c *tcpSession) SendRaw(ctx context.Context, pdu []byte) ([]byte, error) { + if len(pdu) < 1 { + return nil, fmt.Errorf("empty pdu") + } + + request := &gridmodbus.ProtocolDataUnit{ + FunctionCode: pdu[0], + Data: append([]byte(nil), pdu[1:]...), + } + aduRequest, err := c.handler.Encode(request) + if err != nil { + return nil, err + } + aduResponse, err := c.handler.Send(ctx, aduRequest) + if err != nil { + return nil, err + } + if err := c.handler.Verify(aduRequest, aduResponse); err != nil { + return nil, err + } + response, err := c.handler.Decode(aduResponse) + if err != nil { + return nil, err + } + if response.FunctionCode != request.FunctionCode { + exceptionCode := byte(0) + if len(response.Data) > 0 { + exceptionCode = response.Data[0] + } + return nil, &gridmodbus.Error{ + FunctionCode: response.FunctionCode, + ExceptionCode: exceptionCode, + } + } + + out := make([]byte, 1+len(response.Data)) + out[0] = response.FunctionCode + copy(out[1:], response.Data) + return out, nil +} + func (c *tcpSession) BeginRequest(ctx context.Context, attemptTimeout time.Duration) func() error { c.handler.Timeout = attemptTimeout @@ -821,7 +863,13 @@ func ValidateRequest(req *Request) error { return newValidationError(ExcIllegalValue, "write data has %d bytes, expected %d", len(req.Data), expected) } default: - return newValidationError(ExcIllegalFunction, "unsupported function code: 0x%02X", req.FunctionCode) + if len(req.PDU) < 1 { + return newValidationError(ExcIllegalFunction, "missing pdu for function code: 0x%02X", req.FunctionCode) + } + if req.PDU[0] != req.FunctionCode { + return newValidationError(ExcIllegalFunction, "pdu function code 0x%02X does not match request 0x%02X", req.PDU[0], req.FunctionCode) + } + return nil } if uint32(req.Address)+uint32(req.Quantity) > 65536 { return newValidationError(ExcIllegalAddress, "address range exceeds 0xFFFF") @@ -882,7 +930,7 @@ func (c *Client) executeRequest(ctx context.Context, req *Request) ([]byte, erro } return c.buildWriteResponse(req.FunctionCode, req.Address, results), nil default: - return nil, fmt.Errorf("unsupported function code: 0x%02X", req.FunctionCode) + return c.session.SendRaw(ctx, req.PDU) } } diff --git a/internal/modbus/client_test.go b/internal/modbus/client_test.go index 150169a..64b3824 100644 --- a/internal/modbus/client_test.go +++ b/internal/modbus/client_test.go @@ -6,6 +6,7 @@ import ( "encoding/binary" "encoding/json" "errors" + "fmt" "io" "log/slog" "net" @@ -90,6 +91,7 @@ func (c *fakeRequestClient) WriteMultipleRegisters(ctx context.Context, _, _ uin } type fakeSession struct { + client *fakeRequestClient connectErr error connectHook func(context.Context) beginHook func(context.Context) @@ -110,6 +112,12 @@ func (c *fakeSession) Close() error { return nil } func (c *fakeSession) SetSlave(byte) {} +func (c *fakeSession) SendRaw(ctx context.Context, _ []byte) ([]byte, error) { + if c.client == nil { + return nil, errors.New("unsupported function code") + } + return c.client.next(ctx) +} func (c *fakeSession) BeginRequest(ctx context.Context, _ time.Duration) func() error { if c.beginHook != nil { c.beginHook(ctx) @@ -134,6 +142,7 @@ func (s *fakeSessionSet) factory() (clientSession, requestClient) { err = s.connectErrs[index] } return &fakeSession{ + client: s.client, connectErr: err, connectHook: s.connectHook, finishErr: s.finishErr, @@ -154,6 +163,14 @@ func readRequest() *Request { return &Request{SlaveID: 1, FunctionCode: FuncReadHoldingRegisters, Address: 32000, Quantity: 1} } +func vendorRequest() *Request { + return &Request{ + SlaveID: 1, + FunctionCode: 0x41, + PDU: []byte{0x41, 0x00, 0x01, 0x02}, + } +} + func writeRequest() *Request { return &Request{ SlaveID: 1, @@ -872,6 +889,263 @@ func TestValidateRequest_AddressRangeBoundaries(t *testing.T) { } } +func TestValidateRequest_VendorFunctionCode(t *testing.T) { + tests := []struct { + name string + req *Request + code byte + }{ + { + name: "huawei login pdu accepted", + req: &Request{FunctionCode: 0x41, PDU: []byte{0x41, 0x00, 0x01, 0x02}}, + }, + { + name: "missing pdu rejected", + req: &Request{FunctionCode: 0x41}, + code: ExcIllegalFunction, + }, + { + name: "mismatched pdu rejected", + req: &Request{FunctionCode: 0x41, PDU: []byte{0x03, 0x00, 0x01}}, + code: ExcIllegalFunction, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := ValidateRequest(tt.req) + if tt.code == 0 { + if err != nil { + t.Fatalf("expected valid vendor request, got %v", err) + } + return + } + if DownstreamException(err) != tt.code { + t.Fatalf("expected exception 0x%02X, got error %v", tt.code, err) + } + }) + } +} + +func TestClient_VendorFunctionReturnsRawPDU(t *testing.T) { + want := []byte{0x41, 0x00, 0xAA} + client, sessions := newFakeClient([]fakeResult{{data: want}}) + + resp, err := client.Execute(t.Context(), vendorRequest()) + if err != nil { + t.Fatalf("execute: %v", err) + } + if !bytes.Equal(resp, want) { + t.Fatalf("response = % x, want % x", resp, want) + } + if sessions.client.callCount() != 1 { + t.Fatalf("wire calls = %d, want 1", sessions.client.callCount()) + } +} + +func TestClient_VendorFunctionDoesNotRetryTransportFailure(t *testing.T) { + client, sessions := newFakeClient([]fakeResult{ + {err: fakeTimeoutError{}}, + {data: []byte{0x41, 0x00, 0xAA}}, + }) + + _, err := client.Execute(t.Context(), vendorRequest()) + if ErrorKindOf(err) != ErrorTransportTimeout { + t.Fatalf("expected transport timeout, got %v", err) + } + if sessions.client.callCount() != 1 || sessions.sessions.Load() != 1 { + t.Fatalf("vendor function retried, calls=%d sessions=%d", sessions.client.callCount(), sessions.sessions.Load()) + } +} + +func TestClient_LoopbackVendorFunction(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + defer listener.Close() + + wantPDU := []byte{0x41, 0x00, 0x01, 0x02} + wantResp := []byte{0x41, 0x00, 0xAA} + serverErr := make(chan error, 1) + go func() { + conn, acceptErr := listener.Accept() + if acceptErr != nil { + serverErr <- acceptErr + return + } + defer conn.Close() + header := make([]byte, mbapHeaderSize) + if _, readErr := io.ReadFull(conn, header); readErr != nil { + serverErr <- readErr + return + } + pduLen := int(binary.BigEndian.Uint16(header[4:6])) - 1 + pdu := make([]byte, pduLen) + if _, readErr := io.ReadFull(conn, pdu); readErr != nil { + serverErr <- readErr + return + } + if !bytes.Equal(pdu, wantPDU) { + serverErr <- fmt.Errorf("upstream pdu = % x, want % x", pdu, wantPDU) + return + } + response := make([]byte, mbapHeaderSize+len(wantResp)) + copy(response[:2], header[:2]) + binary.BigEndian.PutUint16(response[4:6], uint16(len(wantResp)+1)) + response[6] = header[6] + copy(response[7:], wantResp) + _, writeErr := conn.Write(response) + serverErr <- writeErr + }() + + client := NewClient(listener.Addr().String(), time.Second, 0, 0, slog.New(slog.NewTextHandler(io.Discard, nil))) + t.Cleanup(func() { + if err := client.Close(); err != nil { + t.Errorf("close client: %v", err) + } + }) + resp, err := client.Execute(t.Context(), vendorRequest()) + if err != nil { + t.Fatalf("execute: %v", err) + } + if !bytes.Equal(resp, wantResp) { + t.Fatalf("response = % x, want % x", resp, wantResp) + } + if err := <-serverErr; err != nil { + t.Fatalf("loopback server: %v", err) + } +} + +func TestClient_LoopbackVendorException(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + defer listener.Close() + + serverErr := make(chan error, 1) + go func() { + conn, acceptErr := listener.Accept() + if acceptErr != nil { + serverErr <- acceptErr + return + } + defer conn.Close() + request := make([]byte, 11) + if _, readErr := io.ReadFull(conn, request); readErr != nil { + serverErr <- readErr + return + } + response := make([]byte, 9) + copy(response[:2], request[:2]) + binary.BigEndian.PutUint16(response[4:6], 3) + response[6] = request[6] + response[7] = request[7] | 0x80 + response[8] = ExcIllegalValue + _, writeErr := conn.Write(response) + serverErr <- writeErr + }() + + client := NewClient(listener.Addr().String(), time.Second, 0, 0, slog.New(slog.NewTextHandler(io.Discard, nil))) + t.Cleanup(func() { + if err := client.Close(); err != nil { + t.Errorf("close client: %v", err) + } + }) + _, err = client.Execute(t.Context(), vendorRequest()) + if ErrorKindOf(err) != ErrorProtocolException { + t.Fatalf("expected protocol exception, got %v", err) + } + if DownstreamException(err) != ExcIllegalValue { + t.Fatalf("expected preserved exception, got 0x%02X", DownstreamException(err)) + } + if err := <-serverErr; err != nil { + t.Fatalf("loopback server: %v", err) + } +} + +func TestClient_LoopbackVendorMalformedResponsesMapGatewayWithoutRetry(t *testing.T) { + tests := []struct { + name string + buildResponse func([]byte) []byte + }{ + { + name: "wrong normal function code", + buildResponse: func(request []byte) []byte { + response := make([]byte, 11) + copy(response[:2], request[:2]) + binary.BigEndian.PutUint16(response[4:6], 5) + response[6] = request[6] + response[7] = FuncReadHoldingRegisters + response[8] = 2 + response[9] = 0x12 + response[10] = 0x34 + return response + }, + }, + { + name: "exception missing code", + buildResponse: func(request []byte) []byte { + response := make([]byte, 8) + copy(response[:2], request[:2]) + binary.BigEndian.PutUint16(response[4:6], 2) + response[6] = request[6] + response[7] = request[7] | 0x80 + return response + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + defer listener.Close() + + var accepts atomic.Int32 + serverErr := make(chan error, 1) + go func() { + conn, acceptErr := listener.Accept() + if acceptErr != nil { + serverErr <- acceptErr + return + } + accepts.Add(1) + defer conn.Close() + request := make([]byte, 11) + if _, readErr := io.ReadFull(conn, request); readErr != nil { + serverErr <- readErr + return + } + _, writeErr := conn.Write(tt.buildResponse(request)) + serverErr <- writeErr + }() + + client := NewClient(listener.Addr().String(), time.Second, 0, 0, slog.New(slog.NewTextHandler(io.Discard, nil))) + t.Cleanup(func() { + if err := client.Close(); err != nil { + t.Errorf("close client: %v", err) + } + }) + _, err = client.Execute(t.Context(), vendorRequest()) + if ErrorKindOf(err) != ErrorTransportClosed { + t.Fatalf("expected transport closed, got %v", err) + } + if DownstreamException(err) != ExcGatewayTargetFailed { + t.Fatalf("expected gateway exception, got 0x%02X", DownstreamException(err)) + } + if accepts.Load() != 1 { + t.Fatalf("vendor malformed response retried, accepts=%d", accepts.Load()) + } + if err := <-serverErr; err != nil { + t.Fatalf("loopback server: %v", err) + } + }) + } +} + func TestClient_UnknownUpstreamErrorsRetryReadsAndMapToGatewayFailure(t *testing.T) { client, sessions := newFakeClient([]fakeResult{ {err: errors.New("modbus: transaction id mismatch")}, diff --git a/internal/modbus/server.go b/internal/modbus/server.go index 6a28863..8e9f3a1 100644 --- a/internal/modbus/server.go +++ b/internal/modbus/server.go @@ -48,6 +48,7 @@ type Request struct { Address uint16 Quantity uint16 Data []byte // For write operations + PDU []byte // Raw PDU, including function code } // Handler interface for processing Modbus requests. @@ -302,6 +303,8 @@ func (s *Server) readRequest(conn net.Conn) (*Request, error) { } func (s *Server) parsePDU(req *Request, pdu []byte) error { + req.PDU = append([]byte(nil), pdu...) + switch req.FunctionCode { case FuncReadCoils, FuncReadDiscreteInputs, FuncReadHoldingRegisters, FuncReadInputRegisters: if len(pdu) < 5 { @@ -331,7 +334,7 @@ func (s *Server) parsePDU(req *Request, pdu []byte) error { req.Data = pdu[6 : 6+byteCount] default: - // Unknown function code - let handler deal with it + // Unknown function codes keep the raw PDU for opaque forwarding. } return nil diff --git a/internal/modbus/server_test.go b/internal/modbus/server_test.go index 16b179f..2e7222b 100644 --- a/internal/modbus/server_test.go +++ b/internal/modbus/server_test.go @@ -150,6 +150,11 @@ func TestServer_ParsePDU(t *testing.T) { wantQty: 2, wantData: []byte{0x00, 0x0A, 0x01, 0x02}, }, + { + name: "vendor function keeps raw pdu", + funcCode: 0x41, + pdu: []byte{0x41, 0x00, 0x01, 0x02, 0x03}, + }, } logger := slog.New(slog.NewTextHandler(io.Discard, nil)) @@ -169,6 +174,9 @@ func TestServer_ParsePDU(t *testing.T) { if req.Quantity != tt.wantQty { t.Errorf("quantity: got %d, want %d", req.Quantity, tt.wantQty) } + if !bytes.Equal(req.PDU, tt.pdu) { + t.Errorf("pdu: got % x, want % x", req.PDU, tt.pdu) + } if tt.wantData != nil { if len(req.Data) != len(tt.wantData) { t.Errorf("data length: got %d, want %d", len(req.Data), len(tt.wantData)) @@ -183,6 +191,67 @@ func TestServer_ParsePDU(t *testing.T) { } } +func TestServer_ForwardsVendorPDUToHandler(t *testing.T) { + var got *Request + handler := HandlerFunc(func(_ context.Context, req *Request) ([]byte, error) { + got = req + return []byte{0x41, 0x00, 0xAA}, nil + }) + logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + server := NewServer(handler, time.Second, logger) + if err := server.Listen("127.0.0.1:0"); err != nil { + t.Fatalf("listen: %v", err) + } + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + go func() { + if err := server.Serve(ctx); err != nil { + t.Errorf("serve: %v", err) + } + }() + t.Cleanup(func() { + if err := server.Close(); err != nil { + t.Errorf("close server: %v", err) + } + }) + + conn, err := net.DialTimeout("tcp", server.listener.Addr().String(), time.Second) + if err != nil { + t.Fatalf("dial: %v", err) + } + defer conn.Close() + + pdu := []byte{0x41, 0x00, 0x01, 0x02, 0x03} + frame := make([]byte, 7+len(pdu)) + binary.BigEndian.PutUint16(frame[0:2], 7) + binary.BigEndian.PutUint16(frame[4:6], uint16(len(pdu)+1)) + frame[6] = 1 + copy(frame[7:], pdu) + if _, err := conn.Write(frame); err != nil { + t.Fatalf("write: %v", err) + } + + conn.SetReadDeadline(time.Now().Add(time.Second)) + header := make([]byte, mbapHeaderSize) + if _, err := io.ReadFull(conn, header); err != nil { + t.Fatalf("read header: %v", err) + } + pduLen := int(binary.BigEndian.Uint16(header[4:6])) - 1 + if pduLen < 1 { + t.Fatalf("invalid response length: %d", pduLen) + } + respPDU := make([]byte, pduLen) + if _, err := io.ReadFull(conn, respPDU); err != nil { + t.Fatalf("read pdu: %v", err) + } + if !bytes.Equal(respPDU, []byte{0x41, 0x00, 0xAA}) { + t.Fatalf("unexpected response: % x", respPDU) + } + if got == nil || got.FunctionCode != 0x41 || !bytes.Equal(got.PDU, pdu) { + t.Fatalf("handler request = %+v", got) + } +} + func TestServer_PreservesUpstreamException(t *testing.T) { handler := &mockHandler{err: &RequestError{ Kind: ErrorProtocolException, diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go index b170e75..3aa8ccb 100644 --- a/internal/proxy/proxy.go +++ b/internal/proxy/proxy.go @@ -159,7 +159,15 @@ func (p *Proxy) HandleRequest(ctx context.Context, req *modbus.Request) ([]byte, return p.handleRead(ctx, req) } - return nil, fmt.Errorf("validated unsupported function code: 0x%02X", req.FunctionCode) + return p.handlePassthrough(ctx, req) +} + +func (p *Proxy) handlePassthrough(ctx context.Context, req *modbus.Request) ([]byte, error) { + p.logger.Debug("forwarding function code", + "slave_id", req.SlaveID, + "func", fmt.Sprintf("0x%02X", req.FunctionCode), + ) + return p.client.Execute(ctx, req) } func (p *Proxy) handleRead(ctx context.Context, req *modbus.Request) ([]byte, error) { diff --git a/internal/proxy/proxy_test.go b/internal/proxy/proxy_test.go index bb84346..0c4068f 100644 --- a/internal/proxy/proxy_test.go +++ b/internal/proxy/proxy_test.go @@ -21,6 +21,7 @@ type mockClient struct { response []byte err error calls int + lastReq *modbus.Request } type healthReportingClient struct { @@ -112,6 +113,7 @@ func (m *mockClient) Healthy() error { return nil } func (m *mockClient) Execute(ctx context.Context, req *modbus.Request) ([]byte, error) { m.calls++ + m.lastReq = req return m.response, m.err } @@ -539,37 +541,136 @@ func TestProxy_HandleWriteReadOnlyMode(t *testing.T) { } } -func TestProxy_HandleUnknownFunction(t *testing.T) { +func TestProxy_HandleUnknownFunctionForwardsUpstream(t *testing.T) { logger := slog.New(slog.NewTextHandler(io.Discard, nil)) c := cache.New(time.Second, false) defer c.Close() + upstream := &mockClient{response: []byte{0x41, 0x00, 0xAA}} p := &Proxy{ cfg: &config.Config{ ReadOnly: config.ReadOnlyOn, }, logger: logger, + client: upstream, cache: c, } req := &modbus.Request{ SlaveID: 1, - FunctionCode: 0x99, // Unknown function - Address: 0, - Quantity: 1, + FunctionCode: 0x41, + PDU: []byte{0x41, 0x00, 0x01, 0x02}, } resp, err := p.HandleRequest(context.Background(), req) if err != nil { t.Fatalf("unexpected error: %v", err) } + if !bytes.Equal(resp, upstream.response) { + t.Fatalf("forwarded response: got % x, want % x", resp, upstream.response) + } + if upstream.calls != 1 { + t.Fatalf("upstream calls = %d, want 1", upstream.calls) + } + if upstream.lastReq == nil { + t.Fatal("upstream did not receive the vendor request") + } + if !bytes.Equal(upstream.lastReq.PDU, req.PDU) { + t.Fatalf("upstream request PDU = % x, want % x", upstream.lastReq.PDU, req.PDU) + } +} + +func TestProxy_VendorFunctionForwardsInReadOnlyDeny(t *testing.T) { + upstream := &mockClient{response: []byte{0x41, 0x00, 0xAA}} + p := &Proxy{ + cfg: &config.Config{ReadOnly: config.ReadOnlyDeny}, + logger: slog.New(slog.NewTextHandler(io.Discard, nil)), + client: upstream, + cache: cache.New(time.Second, false), + } + defer p.cache.Close() + + resp, err := p.HandleRequest(context.Background(), &modbus.Request{ + SlaveID: 1, + FunctionCode: 0x41, + PDU: []byte{0x41, 0x00, 0x01}, + }) + if err != nil { + t.Fatalf("vendor request: %v", err) + } + if !bytes.Equal(resp, upstream.response) { + t.Fatalf("deny mode blocked vendor passthrough: % x", resp) + } + if upstream.calls != 1 { + t.Fatalf("upstream calls = %d, want 1", upstream.calls) + } +} + +func TestProxy_VendorFunctionDoesNotUseOrInvalidateCache(t *testing.T) { + logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + c := cache.New(time.Second, false) + defer c.Close() + c.SetRange(1, modbus.FuncReadHoldingRegisters, 0, [][]byte{{0x00, 0x0A}}) + upstream := &mockClient{response: []byte{0x41, 0x00, 0xAA}} + p := &Proxy{ + cfg: &config.Config{ReadOnly: config.ReadOnlyOn}, + logger: logger, + client: upstream, + cache: c, + } - // Should return exception response - if resp[0] != 0x99|0x80 { - t.Errorf("expected exception function code 0x%02X, got 0x%02X", 0x99|0x80, resp[0]) + vendor := &modbus.Request{ + SlaveID: 1, + FunctionCode: 0x41, + PDU: []byte{0x41, 0x00, 0x01}, + } + if _, err := p.HandleRequest(context.Background(), vendor); err != nil { + t.Fatalf("vendor request: %v", err) + } + if upstream.calls != 1 { + t.Fatalf("vendor request calls = %d, want 1", upstream.calls) + } + + read := &modbus.Request{ + SlaveID: 1, + FunctionCode: modbus.FuncReadHoldingRegisters, + Address: 0, + Quantity: 1, + } + resp, err := p.HandleRequest(context.Background(), read) + if err != nil { + t.Fatalf("cached read: %v", err) + } + if !bytes.Equal(resp, []byte{0x03, 0x02, 0x00, 0x0A}) { + t.Fatalf("cached read changed after vendor passthrough: % x", resp) + } + if upstream.calls != 1 { + t.Fatalf("cached read reached upstream, calls = %d", upstream.calls) } - if resp[1] != modbus.ExcIllegalFunction { - t.Errorf("expected illegal function exception, got 0x%02X", resp[1]) +} + +func TestProxy_MissingVendorPDUReturnsIllegalFunctionWithoutUpstream(t *testing.T) { + upstream := &mockClient{} + p := &Proxy{ + cfg: &config.Config{ReadOnly: config.ReadOnlyOff}, + logger: slog.New(slog.NewTextHandler(io.Discard, nil)), + client: upstream, + cache: cache.New(time.Second, false), + } + defer p.cache.Close() + + resp, err := p.HandleRequest(context.Background(), &modbus.Request{ + SlaveID: 1, + FunctionCode: 0x41, + }) + if err != nil { + t.Fatalf("handle request: %v", err) + } + if !bytes.Equal(resp, []byte{0x41 | 0x80, modbus.ExcIllegalFunction}) { + t.Fatalf("unexpected validation response: % x", resp) + } + if upstream.calls != 0 { + t.Fatalf("invalid vendor request reached upstream %d times", upstream.calls) } }