diff --git a/agent_test.go b/agent_test.go index fe2b68b..de9e1f3 100644 --- a/agent_test.go +++ b/agent_test.go @@ -85,8 +85,9 @@ func testPools(t *testing.T) *jobPools { snmp: newWorkerPool(4), mikrotik: newWorkerPool(4), ping: newWorkerPool(4), + checks: newWorkerPool(4), } - t.Cleanup(func() { p.snmp.stop(); p.mikrotik.stop(); p.ping.stop() }) + t.Cleanup(func() { p.snmp.stop(); p.mikrotik.stop(); p.ping.stop(); p.checks.stop() }) return p } @@ -104,7 +105,8 @@ func TestHandleMessage(t *testing.T) { mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) - handleMessage(context.Background(), channelMsg{Event: "phx_reply", Payload: json.RawMessage(`{}`)}, testPools(t), snmpCh, mtCh, credCh, monCh) + checkCh := make(chan *pb.CheckResult, 1) + handleMessage(context.Background(), channelMsg{Event: "phx_reply", Payload: json.RawMessage(`{}`)}, testPools(t), snmpCh, mtCh, credCh, monCh, checkCh) // Just verify it doesn't panic }) @@ -124,6 +126,7 @@ func TestHandleMessage(t *testing.T) { mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) + checkCh := make(chan *pb.CheckResult, 1) payload := makeJobPayload(&pb.AgentJob{ JobId: "j1", @@ -131,7 +134,7 @@ func TestHandleMessage(t *testing.T) { SnmpDevice: &pb.SnmpDevice{Ip: "10.0.0.1", Port: 161}, }) - handleMessage(context.Background(), channelMsg{Event: "jobs", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh) + handleMessage(context.Background(), channelMsg{Event: "jobs", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh, checkCh) // Wait for goroutine to finish select { case <-snmpCh: @@ -145,7 +148,8 @@ func TestHandleMessage(t *testing.T) { mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) - handleMessage(context.Background(), channelMsg{Event: "jobs", Payload: json.RawMessage(`not json`)}, testPools(t), snmpCh, mtCh, credCh, monCh) + checkCh := make(chan *pb.CheckResult, 1) + handleMessage(context.Background(), channelMsg{Event: "jobs", Payload: json.RawMessage(`not json`)}, testPools(t), snmpCh, mtCh, credCh, monCh, checkCh) // Should log error but not panic }) @@ -154,8 +158,9 @@ func TestHandleMessage(t *testing.T) { mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) + checkCh := make(chan *pb.CheckResult, 1) payload, _ := json.Marshal(map[string]string{"binary": "not-base64!!!"}) - handleMessage(context.Background(), channelMsg{Event: "jobs", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh) + handleMessage(context.Background(), channelMsg{Event: "jobs", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh, checkCh) }) t.Run("invalid protobuf", func(t *testing.T) { @@ -163,8 +168,9 @@ func TestHandleMessage(t *testing.T) { mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) + checkCh := make(chan *pb.CheckResult, 1) payload, _ := json.Marshal(map[string]string{"binary": base64.StdEncoding.EncodeToString([]byte{0xFF, 0xFF, 0xFF})}) - handleMessage(context.Background(), channelMsg{Event: "jobs", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh) + handleMessage(context.Background(), channelMsg{Event: "jobs", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh, checkCh) }) t.Run("restart", func(t *testing.T) { @@ -172,7 +178,8 @@ func TestHandleMessage(t *testing.T) { mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) - shouldEnd := handleMessage(context.Background(), channelMsg{Event: "restart", Payload: json.RawMessage(`{}`)}, testPools(t), snmpCh, mtCh, credCh, monCh) + checkCh := make(chan *pb.CheckResult, 1) + shouldEnd := handleMessage(context.Background(), channelMsg{Event: "restart", Payload: json.RawMessage(`{}`)}, testPools(t), snmpCh, mtCh, credCh, monCh, checkCh) if !shouldEnd { t.Error("expected handleMessage to return true for restart") @@ -193,8 +200,9 @@ func TestHandleMessage(t *testing.T) { mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) + checkCh := make(chan *pb.CheckResult, 1) payload, _ := json.Marshal(map[string]string{"url": "https://example.com/agent", "checksum": "abc123"}) - handleMessage(context.Background(), channelMsg{Event: "update", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh) + handleMessage(context.Background(), channelMsg{Event: "update", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh, checkCh) if calledURL != "https://example.com/agent" { t.Errorf("expected update URL %q, got %q", "https://example.com/agent", calledURL) @@ -215,9 +223,10 @@ func TestHandleMessage(t *testing.T) { mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) + checkCh := make(chan *pb.CheckResult, 1) // Missing URL field payload, _ := json.Marshal(map[string]string{"checksum": "abc123"}) - handleMessage(context.Background(), channelMsg{Event: "update", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh) + handleMessage(context.Background(), channelMsg{Event: "update", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh, checkCh) if called { t.Error("selfUpdate should not be called with empty URL") @@ -236,8 +245,9 @@ func TestHandleMessage(t *testing.T) { mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) + checkCh := make(chan *pb.CheckResult, 1) payload, _ := json.Marshal(map[string]string{"url": "https://example.com/agent", "checksum": "abc123"}) - handleMessage(context.Background(), channelMsg{Event: "update", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh) + handleMessage(context.Background(), channelMsg{Event: "update", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh, checkCh) // Should log error but not panic }) @@ -255,8 +265,9 @@ func TestHandleMessage(t *testing.T) { mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) + checkCh := make(chan *pb.CheckResult, 1) payload, _ := json.Marshal(map[string]string{"url": "https://example.com/agent"}) - handleMessage(context.Background(), channelMsg{Event: "update", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh) + handleMessage(context.Background(), channelMsg{Event: "update", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh, checkCh) if called { t.Error("selfUpdate should not be called with empty checksum") @@ -268,7 +279,8 @@ func TestHandleMessage(t *testing.T) { mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) - handleMessage(context.Background(), channelMsg{Event: "some_unknown_event", Payload: json.RawMessage(`{}`)}, testPools(t), snmpCh, mtCh, credCh, monCh) + checkCh := make(chan *pb.CheckResult, 1) + handleMessage(context.Background(), channelMsg{Event: "some_unknown_event", Payload: json.RawMessage(`{}`)}, testPools(t), snmpCh, mtCh, credCh, monCh, checkCh) // Should just log and not panic }) @@ -288,13 +300,14 @@ func TestHandleMessage(t *testing.T) { mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) + checkCh := make(chan *pb.CheckResult, 1) payload := makeJobPayload(&pb.AgentJob{ JobId: "d1", JobType: pb.JobType_DISCOVER, SnmpDevice: &pb.SnmpDevice{Ip: "10.0.0.1"}, }) - handleMessage(context.Background(), channelMsg{Event: "discovery_job", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh) + handleMessage(context.Background(), channelMsg{Event: "discovery_job", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh, checkCh) select { case <-snmpCh: case <-time.After(2 * time.Second): @@ -315,13 +328,14 @@ func TestHandleMessage(t *testing.T) { mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) + checkCh := make(chan *pb.CheckResult, 1) payload := makeJobPayload(&pb.AgentJob{ JobId: "backup:dev1", JobType: pb.JobType_MIKROTIK, MikrotikDevice: &pb.MikrotikDevice{Ip: "10.0.0.1", SshPort: 22, Username: "admin", Password: "pass"}, }) - handleMessage(context.Background(), channelMsg{Event: "backup_job", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh) + handleMessage(context.Background(), channelMsg{Event: "backup_job", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh, checkCh) select { case result := <-mtCh: if result.Error != "" { @@ -338,13 +352,14 @@ func TestHandleMessageRejectsOversizedPayload(t *testing.T) { mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) + checkCh := make(chan *pb.CheckResult, 1) // Create a binary payload larger than maxJobPayloadBytes oversized := make([]byte, maxJobPayloadBytes+1) encoded := base64.StdEncoding.EncodeToString(oversized) payload, _ := json.Marshal(map[string]string{"binary": encoded}) - handleMessage(context.Background(), channelMsg{Event: "jobs", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh) + handleMessage(context.Background(), channelMsg{Event: "jobs", Payload: payload}, testPools(t), snmpCh, mtCh, credCh, monCh, checkCh) // Verify no jobs were dispatched select { @@ -442,12 +457,13 @@ func TestDispatchJob(t *testing.T) { mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) + checkCh := make(chan *pb.CheckResult, 1) dispatchJob(context.Background(), &pb.AgentJob{ JobId: "mt1", JobType: pb.JobType_MIKROTIK, MikrotikDevice: &pb.MikrotikDevice{Ip: "10.0.0.1", Port: 8728}, - }, testPools(t), snmpCh, mtCh, credCh, monCh) + }, testPools(t), snmpCh, mtCh, credCh, monCh, checkCh) select { case result := <-mtCh: @@ -470,12 +486,13 @@ func TestDispatchJob(t *testing.T) { mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) + checkCh := make(chan *pb.CheckResult, 1) dispatchJob(context.Background(), &pb.AgentJob{ JobId: "tc1", JobType: pb.JobType_TEST_CREDENTIALS, SnmpDevice: &pb.SnmpDevice{Ip: "10.0.0.1"}, - }, testPools(t), snmpCh, mtCh, credCh, monCh) + }, testPools(t), snmpCh, mtCh, credCh, monCh, checkCh) select { case result := <-credCh: @@ -498,12 +515,13 @@ func TestDispatchJob(t *testing.T) { mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) + checkCh := make(chan *pb.CheckResult, 1) dispatchJob(context.Background(), &pb.AgentJob{ JobId: "p1", JobType: pb.JobType_PING, SnmpDevice: &pb.SnmpDevice{Ip: "127.0.0.1"}, - }, testPools(t), snmpCh, mtCh, credCh, monCh) + }, testPools(t), snmpCh, mtCh, credCh, monCh, checkCh) select { case result := <-monCh: @@ -530,12 +548,13 @@ func TestDispatchJob(t *testing.T) { mtCh := make(chan *pb.MikrotikResult, 1) credCh := make(chan *pb.CredentialTestResult, 1) monCh := make(chan *pb.MonitoringCheck, 1) + checkCh := make(chan *pb.CheckResult, 1) dispatchJob(context.Background(), &pb.AgentJob{ JobId: "s1", JobType: pb.JobType_POLL, SnmpDevice: &pb.SnmpDevice{Ip: "10.0.0.1"}, - }, testPools(t), snmpCh, mtCh, credCh, monCh) + }, testPools(t), snmpCh, mtCh, credCh, monCh, checkCh) select { case <-snmpCh: diff --git a/checks.go b/checks.go index 2f7bd98..a148158 100644 --- a/checks.go +++ b/checks.go @@ -7,6 +7,7 @@ import ( "net" "net/http" "regexp" + "strconv" "strings" "time" @@ -137,7 +138,7 @@ func executeHTTPCheck(ctx context.Context, config *pb.HttpCheckConfig, timeoutMs // executeTCPCheck performs a TCP port connectivity check func executeTCPCheck(ctx context.Context, config *pb.TcpCheckConfig, timeoutMs uint32) (uint32, string, float64) { timeout := time.Duration(timeoutMs) * time.Millisecond - address := fmt.Sprintf("%s:%d", config.Host, config.Port) + address := net.JoinHostPort(config.Host, strconv.Itoa(int(config.Port))) startTime := time.Now() conn, err := net.DialTimeout("tcp", address, timeout)