Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 13 additions & 0 deletions internal/adapter/input/http/handler_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -243,6 +243,19 @@ func TestDocs_AsyncAPI(t *testing.T) {
}
}

func TestGetMessageByID_Returns501(t *testing.T) {
router := newTestRouter(func(_ context.Context, _ domain.InputType, _ string, _ []byte) (string, error) {
return "", nil
})
req := httptest.NewRequest(http.MethodGet, "/inputs/beszel/messages/some-id", nil)
req.Header.Set("Authorization", "Bearer test-token")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusNotImplemented {
t.Errorf("status = %d, want 501", w.Code)
}
}

func TestDocs_HTML(t *testing.T) {
router := newTestRouter(func(_ context.Context, _ domain.InputType, _ string, _ []byte) (string, error) {
return "", nil
Expand Down
5 changes: 4 additions & 1 deletion internal/adapter/input/http/router.go
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,10 @@ func NewRouter(uc input.ReceiveMessageUseCase, resolver input.InputResolver, ws
ws.ServeWS(w, req, inputTypeFromContext(req.Context()))
})
r.Post("/messages", h.PostMessage)
r.Get("/messages/{messageId}", h.Healthz) // placeholder
r.Get("/messages/{messageId}", func(w http.ResponseWriter, r *http.Request) {
writeError(w, r, http.StatusNotImplemented, "Not Implemented",
"get message by ID is not yet implemented")
})
})

return r
Expand Down
61 changes: 61 additions & 0 deletions internal/adapter/input/tcp/listener_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -206,6 +206,67 @@ func TestListener_GracefulShutdown(t *testing.T) {
}
}

func TestListener_MaxMessageSize_Accepted(t *testing.T) {
// A message body of exactly (maxMessageBytes - 1) bytes (+ delimiter) must be delivered.
mock := &mockReceiveUseCase{returnID: "msg-1"}
addr, _ := startTestListener(t, '\n', "application/json", mock)

conn, err := net.Dial("tcp", addr)
if err != nil {
t.Fatalf("dial: %v", err)
}
defer conn.Close()

// payload = maxMessageBytes - 1 bytes, then delimiter
payload := make([]byte, maxMessageBytes-1)
for i := range payload {
payload[i] = 'x'
}
payload = append(payload, '\n')

if _, err := conn.Write(payload); err != nil {
t.Fatalf("write: %v", err)
}

calls, err := waitForCalls(mock, 1, 3*time.Second)
if err != nil {
t.Fatalf("message not received: %v", err)
}
if len(calls[0].body) != maxMessageBytes-1 {
t.Errorf("body len = %d, want %d", len(calls[0].body), maxMessageBytes-1)
}
}

func TestListener_MaxMessageSize_Exceeded(t *testing.T) {
// A message body exceeding maxMessageBytes must be silently dropped by the scanner.
mock := &mockReceiveUseCase{returnID: "msg-1"}
addr, _ := startTestListener(t, '\n', "application/json", mock)

conn, err := net.Dial("tcp", addr)
if err != nil {
t.Fatalf("dial: %v", err)
}
defer conn.Close()

// payload = maxMessageBytes + 1 bytes, then delimiter — scanner will reject this
payload := make([]byte, maxMessageBytes+1)
for i := range payload {
payload[i] = 'x'
}
payload = append(payload, '\n')

if _, err := conn.Write(payload); err != nil {
t.Fatalf("write: %v", err)
}

// Wait and verify no calls were recorded
time.Sleep(500 * time.Millisecond)
calls := mock.getCalls()
if len(calls) != 0 {
t.Errorf("expected 0 calls for oversized message, got %d", len(calls))
}
}

func TestListener_EmptyMessageSkip(t *testing.T) {
mock := &mockReceiveUseCase{returnID: "msg-1"}
addr, _ := startTestListener(t, '\n', "application/json", mock)
Expand Down
21 changes: 17 additions & 4 deletions internal/adapter/output/webhook/sender.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ import (
"context"
"crypto/tls"
"fmt"
"io"
"net/http"
"time"

Expand All @@ -13,19 +14,30 @@ import (

const defaultTimeoutSec = 10

type Sender struct{}
type Sender struct {
transport *http.Transport
insecureTransport *http.Transport
}

func NewSender() *Sender { return &Sender{} }
func NewSender() *Sender {
return &Sender{
transport: &http.Transport{},
insecureTransport: &http.Transport{
TLSClientConfig: &tls.Config{InsecureSkipVerify: true}, //nolint:gosec
},
}
}

func (s *Sender) Send(ctx context.Context, out domain.Output, payload []byte) error {
timeoutSec := out.TimeoutSec
if timeoutSec <= 0 {
timeoutSec = defaultTimeoutSec
}
client := &http.Client{Timeout: time.Duration(timeoutSec) * time.Second}
t := s.transport
if out.SkipTLSVerify {
client.Transport = &http.Transport{TLSClientConfig: &tls.Config{InsecureSkipVerify: true}} //nolint:gosec
t = s.insecureTransport
}
client := &http.Client{Transport: t, Timeout: time.Duration(timeoutSec) * time.Second}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, out.URL, bytes.NewReader(payload))
if err != nil {
return fmt.Errorf("create request: %w", err)
Expand All @@ -39,6 +51,7 @@ func (s *Sender) Send(ctx context.Context, out domain.Output, payload []byte) er
return fmt.Errorf("send: %w", err)
}
defer resp.Body.Close()
_, _ = io.Copy(io.Discard, resp.Body)
if resp.StatusCode >= 400 {
return fmt.Errorf("webhook returned %d", resp.StatusCode)
}
Expand Down
10 changes: 8 additions & 2 deletions internal/application/service/relay_worker.go
Original file line number Diff line number Diff line change
Expand Up @@ -223,6 +223,10 @@ func (w *RelayWorker) deliver(ctx context.Context, out domain.Output, payload []
return fmt.Errorf("retries exhausted: %w", lastErr)
}

var builtinEvalKeys = map[string]struct{}{
"id": {}, "input": {}, "payload": {}, "createdAt": {}, "status": {},
}

func buildEvalData(msg domain.Message) map[string]any {
data := map[string]any{
"id": msg.ID,
Expand All @@ -231,9 +235,11 @@ func buildEvalData(msg domain.Message) map[string]any {
"createdAt": msg.CreatedAt.Format(time.RFC3339),
"status": string(msg.Status),
}
// Merge ParsedData fields
// Merge ParsedData fields, skipping any key that would overwrite a builtin.
for k, v := range msg.ParsedData {
data[k] = v
if _, reserved := builtinEvalKeys[k]; !reserved {
data[k] = v
}
}
return data
}
Expand Down
35 changes: 35 additions & 0 deletions internal/application/service/relay_worker_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -656,3 +656,38 @@ func TestRelayWorker_InvalidTransition_SkipsUpdate(t *testing.T) {
t.Error("expected ack to be called regardless of invalid transition")
}
}

func TestRelayWorker_ParsedDataDoesNotOverrideBuiltinKeys(t *testing.T) {
// ParsedData contains "input": "HACKED" — if this overwrites the builtin
// "input" key, the filter `data.input == "BESZEL"` will fail and the sender
// will never be called. The fix must protect builtin keys from ParsedData.
msg := domain.Message{
ID: "key-collision",
Input: domain.InputTypeBeszel,
Payload: domain.RawPayload(`{}`),
Status: domain.MessageStatusPending,
Version: 1,
ParsedData: map[string]any{
"input": "HACKED", // must NOT overwrite builtin "input" = "BESZEL"
},
}
queue := &mockMessageQueue{messages: []domain.Message{msg}}
repo := &mockRepo{saveFn: func(_ context.Context, _ domain.Message) error { return nil }}
sender := &mockSender{}
ruleReader := &mockRuleReader{
rule: domain.Rule{InputID: "beszel", Filter: `data.input == "BESZEL"`},
outputs: []domain.Output{{ID: "c1", Type: domain.OutputTypeWebhook}},
}
registry := &mockRegistry{sender: sender}

ctx, cancel := context.WithTimeout(context.Background(), 300*time.Millisecond)
defer cancel()

worker := service.NewRelayWorker(queue, repo, ruleReader, registry, newExprRegistry(), service.DefaultRelayWorkerConfig())
worker.Start(ctx, 1)
time.Sleep(150 * time.Millisecond)

if sender.count.Load() == 0 {
t.Error("builtin key 'input' was overwritten by ParsedData — filter failed when it should have passed")
}
}
7 changes: 0 additions & 7 deletions internal/domain/input_type.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,10 +8,3 @@ const (
InputTypeGeneric InputType = "GENERIC"
)

func (s InputType) IsValid() bool {
switch s {
case InputTypeBeszel, InputTypeDozzle, InputTypeGeneric:
return true
}
return false
}