Compare commits

..

10 Commits

Author SHA1 Message Date
Cédric Verstraeten
c2a07672d9 Merge pull request #331 from kerberos-io/feature/add-encrypted-metadata-field
feature/add-encrypted-metadata-field
2026-09-14 22:10:20 +02:00
cedricve
f64cdfd605 feat: add encrypted metadata field to recording uploads 2026-09-14 19:20:24 +00:00
Cédric Verstraeten
af5d728921 Merge pull request #330 from kerberos-io/feature/add-remote-control-ssh-logging
feature/add-remote-control-ssh-logging
2026-09-14 19:27:45 +02:00
Cédric Verstraeten
b6c8855595 Merge branch 'master' into feature/add-remote-control-ssh-logging 2026-09-14 19:25:11 +02:00
Cédric Verstraeten
23ffb19b4d Merge pull request #329 from kerberos-io/feature/support-mqtt-ssl
feature/support-mqtt-ssl
2026-09-14 19:23:59 +02:00
cedricve
7849f34386 feat: add MQTT connection diagnostics and enhance testing for connection establishment 2026-09-14 17:01:26 +00:00
cedricve
50a77b4591 feat: enhance MQTT error logging and update README for MQTT URI 2026-09-14 14:58:27 +00:00
Cédric Verstraeten
0ed64edb19 Merge pull request #327 from firmwarecostum/patch-3
fix ffmpeg 8.1 2: Update codec context and frame deallocation methods
2026-09-07 14:57:44 +02:00
cedricve
e9ef597442 feat: add remote control SSH logging and session management over MQTT 2026-09-07 08:03:41 +00:00
firmwarecostum
e4703fedc3 Update codec context and frame deallocation methods
Replaced avcodec_close with avcodec_free_context for proper resource management.
2026-09-07 07:14:16 +08:00
14 changed files with 814 additions and 17 deletions

View File

@@ -341,9 +341,10 @@ See [RTSPS and TLS certificates](README-RTSPS-TLS.md) for the complete Bosch UI,
| `AGENT_CAPTURE_PIXEL_CHANGE` | If `CONTINUOUS` set to `false`, the number of pixel require to change before motion triggers. | "150" |
| `AGENT_CAPTURE_FRAGMENTED` | Set the format of the recorded MP4 to fragmented (suitable for HLS). | "false" |
| `AGENT_CAPTURE_FRAGMENTED_DURATION` | If `AGENT_CAPTURE_FRAGMENTED` set to `true`, define the duration (seconds) of a fragment. | "8" |
| `AGENT_MQTT_URI` | An MQTT broker endpoint that is used for bi-directional communication (live view, onvif, etc) | "tcp://mqtt.kerberos.io:1883" |
| `AGENT_MQTT_URI` | MQTT broker endpoint for bi-directional communication. Accepts ActiveMQ `mqtt+ssl://` URLs. | "tcp://mqtt.kerberos.io:1883" |
| `AGENT_MQTT_USERNAME` | Username of the MQTT broker. | "" |
| `AGENT_MQTT_PASSWORD` | Password of the MQTT broker. | "" |
| `AGENT_REMOTE_ACCESS_ENABLED` | Allow encrypted Hub MQTT sessions to stream Agent logs and open an interactive shell. Enable only for trusted deployments. | "false" |
| `AGENT_REALTIME_PROCESSING` | If `AGENT_REALTIME_PROCESSING` set to `true`, the agent will send key frames to the topic | "" |
| `AGENT_REALTIME_PROCESSING_TOPIC` | The topic to which keyframes will be sent in base64 encoded format. | "" |
| `AGENT_STUN_URI` | When using WebRTC, you'll need to provide a STUN server. | "stun:turn-fra1.kerberos.io:3478"|
@@ -379,6 +380,13 @@ See [RTSPS and TLS certificates](README-RTSPS-TLS.md) for the complete Bosch UI,
| `AGENT_SIGNING` | Enable 'true' or disable 'false' for signing recordings. | "true" |
| `AGENT_SIGNING_PRIVATE_KEY` | The private key (RSA) to sign the recordings fingerprint to validate origin. | "" - uses default one if empty |
Remote console access is disabled unless `AGENT_REMOTE_ACCESS_ENABLED=true`.
The Agent also rejects remote session messages unless Hub encryption or
end-to-end MQTT encryption is configured and used. A remote shell runs inside
the Agent process environment as the Agent operating-system user; it is not an
SSH server and does not expose a new network port. Keep the feature disabled on
deployments where Hub owners should not have operating-system access.
### Resumable upload chunk size
Hub and Vault resumable uploads use `AGENT_TUS_CHUNK_SIZE_BYTES` as the maximum

View File

@@ -11,6 +11,7 @@ require (
github.com/bluenviron/gortsplib/v5 v5.6.3
github.com/bluenviron/mediacommon v1.14.0
github.com/cedricve/go-onvif v0.0.0-20200222191200-567e8ce298f6
github.com/creack/pty v1.1.24
github.com/dromara/carbon/v2 v2.6.8
github.com/dropbox/dropbox-sdk-go-unofficial/v6 v6.0.5
github.com/eclipse/paho.mqtt.golang v1.5.0

View File

@@ -456,6 +456,8 @@ github.com/cncf/xds/go v0.0.0-20240905190251-b4127c9b8d78/go.mod h1:W+zGtBO5Y1Ig
github.com/cncf/xds/go v0.0.0-20241223141626-cff3c89139a3/go.mod h1:W+zGtBO5Y1IgJhy4+A9GOqVhqLpfZi+vwmdNXUehLA8=
github.com/cncf/xds/go v0.0.0-20250121191232-2f005788dc42/go.mod h1:W+zGtBO5Y1IgJhy4+A9GOqVhqLpfZi+vwmdNXUehLA8=
github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E=
github.com/creack/pty v1.1.24 h1:bJrF4RRfyJnbTJqzRLHzcGaZK1NeM5kTC9jGgovnR1s=
github.com/creack/pty v1.1.24/go.mod h1:08sCNb52WyoAwi2QDyzUCTgcvVFhUzewun7wtTfvcwE=
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=

View File

@@ -1527,13 +1527,13 @@ func newDecoder(codecName string) (*Decoder, error) {
res := C.avcodec_open2(codecCtx, codec, nil)
if res < 0 {
C.avcodec_close(codecCtx)
C.avcodec_free_context(&codecCtx)
return nil, fmt.Errorf("avcodec_open2() failed")
}
srcFrame := C.av_frame_alloc()
if srcFrame == nil {
C.avcodec_close(codecCtx)
C.avcodec_free_context(&codecCtx)
return nil, fmt.Errorf("av_frame_alloc() failed")
}
@@ -1549,7 +1549,7 @@ func (d *Decoder) Close() {
C.av_frame_free(&d.srcFrame)
}
C.av_frame_free(&d.srcFrame)
C.avcodec_close(d.codecCtx)
C.avcodec_free_context(&d.codecCtx)
}
func (d *Decoder) decode(nalu []byte) (image.YCbCr, error) {

View File

@@ -53,12 +53,13 @@ func publishRecordingState(mqttClient mqtt.Client, hubKey string, configuration
}
}
func recordingUploadMetadata(name, deviceKey string, timestamp int64, mp4Video *video.MP4) models.RecordingUploadMetadata {
func recordingUploadMetadata(name, deviceKey string, timestamp int64, mp4Video *video.MP4, encrypted bool) models.RecordingUploadMetadata {
metadata := models.RecordingUploadMetadata{
FileName: filepath.Base(name),
DeviceKey: deviceKey,
Timestamp: timestamp,
Duration: mp4Video.VideoTotalDuration,
Encrypted: encrypted,
}
value := mp4Video.AverageFPS()
if value > 0 && value <= 240 && !math.IsInf(value, 0) && !math.IsNaN(value) {
@@ -481,6 +482,7 @@ func HandleRecordStream(queue *packets.Queue, configDirectory string, configurat
log.Info("capture.main.HandleRecordStream(continuous): no video data recorded, not renaming file.")
}
encrypted := false
// Check if we need to encrypt the recording.
if config.Encryption != nil && config.Encryption.Enabled == "true" && config.Encryption.Recordings == "true" && config.Encryption.SymmetricKey != "" {
// reopen file into memory 'fullName'
@@ -493,6 +495,8 @@ func HandleRecordStream(queue *packets.Queue, configDirectory string, configurat
err := os.WriteFile(fullName, []byte(encryptedContents), 0644)
if err != nil {
log.Error("capture.main.HandleRecordStream(continuous): error writing file: " + err.Error())
} else {
encrypted = true
}
} else {
log.Error("capture.main.HandleRecordStream(continuous): error encrypting file: " + err.Error())
@@ -502,7 +506,7 @@ func HandleRecordStream(queue *packets.Queue, configDirectory string, configurat
}
}
queueRecordingForUpload(configDirectory, recordingUploadMetadata(name, config.Key, startRecording, mp4Video))
queueRecordingForUpload(configDirectory, recordingUploadMetadata(name, config.Key, startRecording, mp4Video, encrypted))
recordingStatus = "idle"
@@ -638,6 +642,7 @@ func HandleRecordStream(queue *packets.Queue, configDirectory string, configurat
log.Info("capture.main.HandleRecordStream(continuous): no video data recorded, not renaming file.")
}
encrypted := false
// Check if we need to encrypt the recording.
if config.Encryption != nil && config.Encryption.Enabled == "true" && config.Encryption.Recordings == "true" && config.Encryption.SymmetricKey != "" {
// reopen file into memory 'fullName'
@@ -650,6 +655,8 @@ func HandleRecordStream(queue *packets.Queue, configDirectory string, configurat
err := os.WriteFile(fullName, []byte(encryptedContents), 0644)
if err != nil {
log.Error("capture.main.HandleRecordStream(motiondetection): error writing file: " + err.Error())
} else {
encrypted = true
}
} else {
log.Error("capture.main.HandleRecordStream(motiondetection): error encrypting file: " + err.Error())
@@ -659,7 +666,7 @@ func HandleRecordStream(queue *packets.Queue, configDirectory string, configurat
}
}
queueRecordingForUpload(configDirectory, recordingUploadMetadata(name, config.Key, startRecording, mp4Video))
queueRecordingForUpload(configDirectory, recordingUploadMetadata(name, config.Key, startRecording, mp4Video, encrypted))
recordingStatus = "idle"
@@ -905,6 +912,7 @@ func HandleRecordStream(queue *packets.Queue, configDirectory string, configurat
log.Info("capture.main.HandleRecordStream(motiondetection): no video data recorded, not renaming file.")
}
encrypted := false
// Check if we need to encrypt the recording.
if config.Encryption != nil && config.Encryption.Enabled == "true" && config.Encryption.Recordings == "true" && config.Encryption.SymmetricKey != "" {
// reopen file into memory 'fullName'
@@ -917,6 +925,8 @@ func HandleRecordStream(queue *packets.Queue, configDirectory string, configurat
err := os.WriteFile(fullName, []byte(encryptedContents), 0644)
if err != nil {
log.Error("capture.main.HandleRecordStream(motiondetection): error writing file: " + err.Error())
} else {
encrypted = true
}
} else {
log.Error("capture.main.HandleRecordStream(motiondetection): error encrypting file: " + err.Error())
@@ -926,7 +936,7 @@ func HandleRecordStream(queue *packets.Queue, configDirectory string, configurat
}
}
queueRecordingForUpload(configDirectory, recordingUploadMetadata(name, config.Key, displayTime, mp4Video))
queueRecordingForUpload(configDirectory, recordingUploadMetadata(name, config.Key, displayTime, mp4Video, encrypted))
// Clean up the recording directory if necessary.
CleanupRecordingDirectory(configDirectory, configuration)

View File

@@ -42,7 +42,7 @@ func TestQueueRecordingForUploadStoresFinalizedMetadata(t *testing.T) {
}
mp4Video := &video.MP4{VideoTotalDuration: 20452, SampleCount: 613}
metadata := recordingUploadMetadata("recording.mp4", "device-key", 1785934709414, mp4Video)
metadata := recordingUploadMetadata("recording.mp4", "device-key", 1785934709414, mp4Video, true)
queueRecordingForUpload(configDirectory, metadata)
got, err := os.ReadFile(filepath.Join(configDirectory, "data", "cloud", "recording.metadata"))
@@ -54,7 +54,7 @@ func TestQueueRecordingForUploadStoresFinalizedMetadata(t *testing.T) {
t.Fatalf("decode upload marker: %v", err)
}
expectedFPS := mp4Video.AverageFPS()
if stored.FileName != "recording.mp4" || stored.DeviceKey != "device-key" || stored.Timestamp != 1785934709414 || stored.Duration != 20452 || math.Abs(stored.FPS-expectedFPS) > 1e-9 {
if stored.FileName != "recording.mp4" || stored.DeviceKey != "device-key" || stored.Timestamp != 1785934709414 || stored.Duration != 20452 || math.Abs(stored.FPS-expectedFPS) > 1e-9 || !stored.Encrypted {
t.Fatalf("upload marker = %+v", stored)
}
if stored.FPS == math.Floor(stored.FPS) {

View File

@@ -15,6 +15,7 @@ import (
const recordingFPSHeader = "X-Kerberos-Storage-Fps"
const recordingDurationHeader = "X-Kerberos-Storage-Duration"
const recordingTimestampHeader = "X-Kerberos-Storage-Timestamp"
const recordingEncryptedHeader = "X-Kerberos-Storage-Encrypted"
// queuedRecordingFPS reads the FPS snapshot written into the upload marker
// when the recording was finalized. Historical empty markers intentionally
@@ -84,5 +85,8 @@ func setQueuedRecordingMetadataHeaders(header http.Header, fileName string) {
if metadata.Timestamp > 0 {
header.Set(recordingTimestampHeader, strconv.FormatInt(metadata.Timestamp, 10))
}
if metadata.Encrypted {
header.Set(recordingEncryptedHeader, "true")
}
}
}

View File

@@ -361,6 +361,9 @@ func addRecordingTusMetadata(values map[string]string, fileName string) {
if metadata.Timestamp > 0 {
values["timestamp"] = strconv.FormatInt(metadata.Timestamp, 10)
}
if metadata.Encrypted {
values["encrypted"] = "true"
}
}
// tusCreate performs the tus "creation" request (POST). On success it returns

View File

@@ -386,7 +386,7 @@ func TestQueuedRecordingFPSAllowsMissingHistoricalMarker(t *testing.T) {
func TestQueuedRecordingMetadataHeaders(t *testing.T) {
fileName := "recording.mp4"
withRecording(t, fileName, []byte("recording"))
withQueuedRecordingFPS(t, fileName, `{"filename":"recording.mp4","device_key":"device-key","timestamp":1785934709414,"duration":20452,"fps":25}`)
withQueuedRecordingFPS(t, fileName, `{"filename":"recording.mp4","device_key":"device-key","timestamp":1785934709414,"duration":20452,"fps":25,"encrypted":true}`)
header := make(http.Header)
setQueuedRecordingMetadataHeaders(header, fileName)
@@ -399,6 +399,14 @@ func TestQueuedRecordingMetadataHeaders(t *testing.T) {
if got := header.Get(recordingTimestampHeader); got != "1785934709414" {
t.Fatalf("timestamp header = %q", got)
}
if got := header.Get(recordingEncryptedHeader); got != "true" {
t.Fatalf("encrypted header = %q", got)
}
metadata := map[string]string{}
addRecordingTusMetadata(metadata, fileName)
if got := metadata["encrypted"]; got != "true" {
t.Fatalf("encrypted TUS metadata = %q", got)
}
}
func TestQueuedRecordingFPSAllowsLegacyMarkerFileName(t *testing.T) {

View File

@@ -324,3 +324,23 @@ type TriggerRelay struct {
DeviceId string `json:"device_id"` // device id
Token string `json:"token"` // token
}
// RemoteSessionPayload controls an interactive shell or log stream over MQTT.
// Data is base64 encoded so terminal control bytes remain valid JSON.
type RemoteSessionPayload struct {
Timestamp int64 `json:"timestamp"`
SessionID string `json:"session_id"`
Kind string `json:"kind,omitempty"`
Data string `json:"data,omitempty"`
Rows uint16 `json:"rows,omitempty"`
Columns uint16 `json:"columns,omitempty"`
Tail int `json:"tail,omitempty"`
}
type RemoteSessionStatus struct {
Timestamp int64 `json:"timestamp"`
SessionID string `json:"session_id"`
Kind string `json:"kind"`
State string `json:"state"`
Error string `json:"error,omitempty"`
}

View File

@@ -16,6 +16,7 @@ type RecordingUploadMetadata struct {
Timestamp int64 `json:"timestamp"` // Unix milliseconds.
Duration uint64 `json:"duration"` // Milliseconds.
FPS float64 `json:"fps,omitempty"`
Encrypted bool `json:"encrypted,omitempty"`
}
// RecordingUploadMetadataFileName returns the queue marker name associated

View File

@@ -1,7 +1,9 @@
package mqtt
import (
"context"
"crypto/rsa"
"crypto/tls"
"crypto/x509"
"encoding/base64"
"encoding/json"
@@ -9,13 +11,14 @@ import (
"fmt"
"io/ioutil"
"math/rand"
"net"
"net/url"
"os"
"strconv"
"strings"
"sync"
"time"
"context"
mqtt "github.com/eclipse/paho.mqtt.golang"
"github.com/kerberos-io/agent/machinery/src/capture"
configService "github.com/kerberos-io/agent/machinery/src/config"
@@ -24,6 +27,7 @@ import (
"github.com/kerberos-io/agent/machinery/src/onvif"
"github.com/kerberos-io/agent/machinery/src/webrtc"
log "github.com/sirupsen/logrus"
"golang.org/x/net/proxy"
)
// We'll cache the MQTT settings to know if we need to reinitialize the MQTT client connection.
@@ -34,6 +38,177 @@ var PREV_MQTTPassword string
var PREV_HubKey string
var PREV_AgentKey string
type pahoErrorLogger struct{}
func (pahoErrorLogger) Println(values ...interface{}) {
log.WithFields(log.Fields{
"component": "routers/mqtt",
"event": "paho_error",
}).Error(strings.TrimSpace(fmt.Sprintln(values...)))
}
func (pahoErrorLogger) Printf(format string, values ...interface{}) {
log.WithFields(log.Fields{
"component": "routers/mqtt",
"event": "paho_error",
}).Errorf(strings.TrimSpace(format), values...)
}
func init() {
mqtt.ERROR = pahoErrorLogger{}
}
func enableMQTTConnectionDiagnostics(options *mqtt.ClientOptions, brokerURL string) {
if !strings.Contains(brokerURL, "://") {
options.SetCustomOpenConnectionFn(openMQTTConnection)
return
}
parsedURL, err := url.Parse(brokerURL)
if err != nil {
return
}
switch strings.ToLower(parsedURL.Scheme) {
case "", "mqtt", "tcp", "ssl", "tls", "mqtts", "mqtt+ssl", "tcps":
options.SetCustomOpenConnectionFn(openMQTTConnection)
}
}
func openMQTTConnection(uri *url.URL, options mqtt.ClientOptions) (net.Conn, error) {
host := uri.Hostname()
fields := log.Fields{
"component": "routers/mqtt",
"host": host,
"port": uri.Port(),
"scheme": uri.Scheme,
}
logMQTTDNSResolution(host, options.ConnectTimeout, fields)
connectionStartedAt := time.Now()
dialer := options.Dialer
if dialer == nil {
dialer = &net.Dialer{Timeout: options.ConnectTimeout}
}
proxyConfigured := os.Getenv("all_proxy") != ""
proxyMode := "direct"
if proxyConfigured {
proxyMode = "socks"
}
fields["proxy_mode"] = proxyMode
log.WithFields(fields).Info("Opening MQTT TCP connection")
var (
connection net.Conn
err error
)
if proxyConfigured {
connection, err = proxy.FromEnvironment().Dial("tcp", uri.Host)
} else {
connection, err = dialer.Dial("tcp", uri.Host)
}
fields["duration_ms"] = time.Since(connectionStartedAt).Milliseconds()
if err != nil {
logMQTTNetworkError("MQTT TCP connection failed", err, fields)
return nil, err
}
fields["local_address"] = connection.LocalAddr().String()
fields["remote_address"] = connection.RemoteAddr().String()
log.WithFields(fields).Info("MQTT TCP connection established")
if !isSecureMQTTScheme(uri.Scheme) {
return connection, nil
}
tlsConfig := options.TLSConfig
if tlsConfig == nil {
tlsConfig = &tls.Config{}
} else {
tlsConfig = tlsConfig.Clone()
}
if tlsConfig.ServerName == "" {
tlsConfig.ServerName = host
}
tlsConnection := tls.Client(connection, tlsConfig)
tlsStartedAt := time.Now()
if options.ConnectTimeout > 0 {
_ = tlsConnection.SetDeadline(connectionStartedAt.Add(options.ConnectTimeout))
}
if err = tlsConnection.Handshake(); err != nil {
_ = connection.Close()
fields["duration_ms"] = time.Since(tlsStartedAt).Milliseconds()
fields["server_name"] = tlsConfig.ServerName
logMQTTNetworkError("MQTT TLS handshake failed", err, fields)
return nil, err
}
_ = tlsConnection.SetDeadline(time.Time{})
state := tlsConnection.ConnectionState()
fields["cipher_suite"] = tls.CipherSuiteName(state.CipherSuite)
fields["duration_ms"] = time.Since(tlsStartedAt).Milliseconds()
fields["server_name"] = tlsConfig.ServerName
fields["tls_version"] = tls.VersionName(state.Version)
log.WithFields(fields).Info("MQTT TLS handshake established")
return tlsConnection, nil
}
func logMQTTDNSResolution(host string, timeout time.Duration, fields log.Fields) {
if host == "" || net.ParseIP(host) != nil {
return
}
lookupTimeout := timeout
if lookupTimeout <= 0 || lookupTimeout > 5*time.Second {
lookupTimeout = 5 * time.Second
}
ctx, cancel := context.WithTimeout(context.Background(), lookupTimeout)
defer cancel()
startedAt := time.Now()
addresses, err := net.DefaultResolver.LookupHost(ctx, host)
dnsFields := cloneLogFields(fields)
dnsFields["duration_ms"] = time.Since(startedAt).Milliseconds()
if err != nil {
logMQTTNetworkError("MQTT broker DNS resolution failed", err, dnsFields)
return
}
dnsFields["resolved_addresses"] = addresses
log.WithFields(dnsFields).Info("MQTT broker DNS resolved")
}
func logMQTTNetworkError(message string, err error, fields log.Fields) {
errorFields := cloneLogFields(fields)
if networkError, ok := err.(net.Error); ok {
errorFields["network_timeout"] = networkError.Timeout()
}
if operationError, ok := err.(*net.OpError); ok {
errorFields["network"] = operationError.Net
errorFields["operation"] = operationError.Op
}
log.WithError(err).WithFields(errorFields).Error(message)
}
func cloneLogFields(fields log.Fields) log.Fields {
cloned := make(log.Fields, len(fields))
for key, value := range fields {
cloned[key] = value
}
return cloned
}
func isSecureMQTTScheme(scheme string) bool {
switch strings.ToLower(scheme) {
case "ssl", "tls", "mqtts", "mqtt+ssl", "tcps":
return true
default:
return false
}
}
func HasMQTTClientModified(configuration *models.Configuration) bool {
MTTURI := configuration.Config.MQTTURI
MTTUsername := configuration.Config.MQTTUsername
@@ -59,6 +234,7 @@ func HasMQTTClientModified(configuration *models.Configuration) bool {
// - kerberos/{hubkey}/device/{devicekey}/motion: a motion signal
func ConfigureMQTT(configDirectory string, configuration *models.Configuration, communication *models.Communication) mqtt.Client {
installRemoteAccessHook()
config := configuration.Config
@@ -116,6 +292,7 @@ func ConfigureMQTT(configDirectory string, configuration *models.Configuration,
// Some extra options to make sure the connection behaves
// properly. More information here: github.com/eclipse/paho.mqtt.golang.
//opts.SetCleanSession(true)
enableMQTTConnectionDiagnostics(opts, mqttURL)
opts.SetCleanSession(false)
opts.SetResumeSubs(true)
opts.SetStore(mqtt.NewMemoryStore())
@@ -203,7 +380,7 @@ func ConfigureMQTT(configDirectory string, configuration *models.Configuration,
"component": "routers/mqtt",
"event": "initial_connection_timeout",
"timeout_ms": (30 * time.Second).Milliseconds(),
}).Error("Timed out establishing initial MQTT connection")
}).Warn("Initial MQTT connection is still retrying")
}
return mqc
}
@@ -273,6 +450,7 @@ func MQTTListenerHandler(mqttClient mqtt.Client, hubKey string, configDirectory
// We will receive all messages from our hub, so we'll need to filter to the relevant device.
if message.Mid != "" && message.Timestamp != 0 && message.DeviceId == configuration.Config.Key {
var payload models.Payload
remoteAuthenticated := false
// Messages might be hidden, if so we'll need to decrypt them using the Kerberos Hub private key.
if message.Hidden && configuration.Config.HubEncryption == "true" {
@@ -289,8 +467,10 @@ func MQTTListenerHandler(mqttClient mqtt.Client, hubKey string, configDirectory
log.Error("routers.mqtt.main.MQTTListenerHandler(): error decrypting message: " + err.Error())
return
}
json.Unmarshal(visibleValue, &payload)
message.Payload = payload
if err := json.Unmarshal(visibleValue, &payload); err == nil {
message.Payload = payload
remoteAuthenticated = true
}
} else {
log.Error("routers.mqtt.main.MQTTListenerHandler(): error decrypting message, no private key provided.")
}
@@ -338,7 +518,9 @@ func MQTTListenerHandler(mqttClient mqtt.Client, hubKey string, configDirectory
log.Error("routers.mqtt.main.MQTTListenerHandler(): error decrypting message: " + err.Error())
return
}
json.Unmarshal(decryptedValue, &payload)
if err := json.Unmarshal(decryptedValue, &payload); err == nil {
remoteAuthenticated = true
}
} else {
log.Error("routers.mqtt.main.MQTTListenerHandler(): error decrypting message, assymetric keys do not match.")
return
@@ -396,6 +578,14 @@ func MQTTListenerHandler(mqttClient mqtt.Client, hubKey string, configDirectory
go HandleReceiveHDCandidates(mqttClient, hubKey, payload, configuration, communication)
case "trigger-relay":
go HandleTriggerRelay(mqttClient, hubKey, payload, configuration, communication)
case "remote-session-open":
go HandleRemoteSessionOpen(mqttClient, hubKey, payload, remoteAuthenticated, configuration)
case "remote-session-input":
go HandleRemoteSessionInput(mqttClient, hubKey, payload, remoteAuthenticated, configuration)
case "remote-session-resize":
go HandleRemoteSessionResize(payload, remoteAuthenticated)
case "remote-session-close":
go HandleRemoteSessionClose(payload, remoteAuthenticated)
}
}

View File

@@ -1,12 +1,95 @@
package mqtt
import (
"crypto/tls"
"errors"
"net/http/httptest"
"net/url"
"os"
"testing"
"time"
mqtt "github.com/eclipse/paho.mqtt.golang"
"github.com/kerberos-io/agent/machinery/src/models"
)
func TestEnableMQTTConnectionDiagnostics(t *testing.T) {
tests := []struct {
name string
brokerURL string
want bool
}{
{name: "ActiveMQ TLS", brokerURL: "mqtt+ssl://broker.example:8883", want: true},
{name: "TCP", brokerURL: "tcp://broker.example:1883", want: true},
{name: "default TCP", brokerURL: "broker.example:1883", want: true},
{name: "WebSocket", brokerURL: "wss://broker.example/mqtt", want: false},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
options := mqtt.NewClientOptions()
enableMQTTConnectionDiagnostics(options, test.brokerURL)
if got := options.CustomOpenConnectionFn != nil; got != test.want {
t.Fatalf("CustomOpenConnectionFn configured = %t, want %t", got, test.want)
}
})
}
}
func TestOpenMQTTConnectionReturnsTCPError(t *testing.T) {
options := *mqtt.NewClientOptions().SetConnectTimeout(100 * time.Millisecond)
brokerURL, err := url.Parse("tcp://127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
connection, err := openMQTTConnection(brokerURL, options)
if connection != nil {
connection.Close()
t.Fatal("openMQTTConnection() returned a connection for an unavailable endpoint")
}
if err == nil {
t.Fatal("openMQTTConnection() returned no TCP error")
}
}
func TestOpenMQTTConnectionEstablishesActiveMQTLS(t *testing.T) {
server := httptest.NewTLSServer(nil)
defer server.Close()
serverURL, err := url.Parse(server.URL)
if err != nil {
t.Fatal(err)
}
brokerURL, err := url.Parse("mqtt+ssl://" + serverURL.Host)
if err != nil {
t.Fatal(err)
}
options := *mqtt.NewClientOptions().
SetConnectTimeout(time.Second).
SetTLSConfig(&tls.Config{InsecureSkipVerify: true}) // #nosec G402 -- local test server
connection, err := openMQTTConnection(brokerURL, options)
if err != nil {
t.Fatalf("openMQTTConnection() error = %v", err)
}
defer connection.Close()
if _, ok := connection.(*tls.Conn); !ok {
t.Fatalf("openMQTTConnection() connection type = %T, want *tls.Conn", connection)
}
}
func TestIsSecureMQTTScheme(t *testing.T) {
for _, scheme := range []string{"ssl", "tls", "mqtts", "mqtt+ssl", "tcps"} {
if !isSecureMQTTScheme(scheme) {
t.Errorf("isSecureMQTTScheme(%q) = false, want true", scheme)
}
}
if isSecureMQTTScheme("tcp") {
t.Error("isSecureMQTTScheme(\"tcp\") = true, want false")
}
}
func TestConfigureMQTTRequiresHubKey(t *testing.T) {
configuration := &models.Configuration{Config: models.Config{Key: "agent-key"}}
@@ -53,3 +136,73 @@ func TestEnqueueLatestAudioDoesNotBlockNilChannel(t *testing.T) {
t.Fatal("enqueueLatestAudio() blocked on a nil channel")
}
}
func TestRemoteAccessRequiresExplicitOptIn(t *testing.T) {
previous, present := os.LookupEnv(remoteAccessEnvironment)
t.Cleanup(func() {
if present {
_ = os.Setenv(remoteAccessEnvironment, previous)
} else {
_ = os.Unsetenv(remoteAccessEnvironment)
}
})
_ = os.Unsetenv(remoteAccessEnvironment)
if remoteAccessEnabled() {
t.Fatal("remoteAccessEnabled() = true without opt-in")
}
_ = os.Setenv(remoteAccessEnvironment, "true")
if !remoteAccessEnabled() {
t.Fatal("remoteAccessEnabled() = false after opt-in")
}
}
func TestNormalizeTerminalSize(t *testing.T) {
rows, columns := normalizeTerminalSize(0, 0)
if rows != 24 || columns != 80 {
t.Fatalf("normalizeTerminalSize(0, 0) = (%d, %d), want (24, 80)", rows, columns)
}
rows, columns = normalizeTerminalSize(500, 500)
if rows != 200 || columns != 400 {
t.Fatalf("normalizeTerminalSize(500, 500) = (%d, %d), want (200, 400)", rows, columns)
}
}
func TestDecodeRemotePayloadRejectsMissingSession(t *testing.T) {
_, err := decodeRemotePayload(models.Payload{Value: map[string]interface{}{
"kind": "shell",
}})
if err == nil {
t.Fatal("decodeRemotePayload() accepted a missing session id")
}
}
func TestRemoteSessionOpenRejectsUnprovenEncryption(t *testing.T) {
previous, present := os.LookupEnv(remoteAccessEnvironment)
t.Cleanup(func() {
if present {
_ = os.Setenv(remoteAccessEnvironment, previous)
} else {
_ = os.Unsetenv(remoteAccessEnvironment)
}
})
_ = os.Setenv(remoteAccessEnvironment, "true")
// The listener passes false when an envelope merely claims to be hidden but
// no ciphertext was successfully decrypted. The remote handler must reject it.
if err := validateRemoteAccess(false); !errors.Is(err, errRemoteUnauthenticated) {
t.Fatalf("validateRemoteAccess(false) error = %v, want %v", err, errRemoteUnauthenticated)
}
}
func TestRemoteSessionReservationIsIdempotent(t *testing.T) {
manager := newRemoteAccessManager()
session := &remoteSession{id: "session-1", kind: "logs"}
if _, err := manager.reserve(session); err != nil {
t.Fatalf("first reserve() failed: %v", err)
}
if _, err := manager.reserve(session); !errors.Is(err, errRemoteSessionExists) {
t.Fatalf("duplicate reserve() error = %v, want %v", err, errRemoteSessionExists)
}
}

View File

@@ -0,0 +1,397 @@
package mqtt
import (
"context"
"encoding/base64"
"encoding/json"
"errors"
"io"
"os"
"os/exec"
"strconv"
"strings"
"sync"
"time"
"github.com/creack/pty"
paho "github.com/eclipse/paho.mqtt.golang"
"github.com/kerberos-io/agent/machinery/src/models"
log "github.com/sirupsen/logrus"
)
const (
remoteAccessEnvironment = "AGENT_REMOTE_ACCESS_ENABLED"
remoteHistoryLimit = 500
remoteSessionLimit = 5
remoteOutputChunkSize = 4096
remoteInputLimit = 64 * 1024
remoteSessionLifetime = time.Hour
)
type remoteSession struct {
id string
kind string
pty *os.File
cancel context.CancelFunc
logs chan string
done chan struct{}
timer *time.Timer
}
type remoteAccessManager struct {
mu sync.Mutex
sessions map[string]*remoteSession
history []string
}
var (
remoteAccess = newRemoteAccessManager()
remoteHookOnce sync.Once
errRemoteDisabled = errors.New("remote access is disabled on this agent")
errRemoteUnauthenticated = errors.New("remote access requires encrypted MQTT")
errRemoteSessionExists = errors.New("remote session already exists")
errRemoteSessionLimit = errors.New("remote session limit reached")
)
func newRemoteAccessManager() *remoteAccessManager {
return &remoteAccessManager{sessions: make(map[string]*remoteSession)}
}
func installRemoteAccessHook() {
remoteHookOnce.Do(func() {
log.AddHook(remoteAccess)
})
}
func (manager *remoteAccessManager) Levels() []log.Level {
return log.AllLevels
}
func (manager *remoteAccessManager) Fire(entry *log.Entry) error {
line, err := json.Marshal(map[string]interface{}{
"timestamp": entry.Time.Format(time.RFC3339Nano),
"level": entry.Level.String(),
"message": entry.Message,
"fields": entry.Data,
})
if err != nil {
return nil
}
encoded := base64.StdEncoding.EncodeToString(append(line, '\n'))
manager.mu.Lock()
manager.history = append(manager.history, encoded)
if len(manager.history) > remoteHistoryLimit {
manager.history = manager.history[len(manager.history)-remoteHistoryLimit:]
}
for _, session := range manager.sessions {
if session.kind != "logs" {
continue
}
select {
case session.logs <- encoded:
default:
}
}
manager.mu.Unlock()
return nil
}
func remoteAccessEnabled() bool {
enabled, err := strconv.ParseBool(strings.TrimSpace(os.Getenv(remoteAccessEnvironment)))
return err == nil && enabled
}
func validateRemoteAccess(authenticated bool) error {
if !authenticated {
return errRemoteUnauthenticated
}
if !remoteAccessEnabled() {
return errRemoteDisabled
}
return nil
}
func decodeRemotePayload(payload models.Payload) (models.RemoteSessionPayload, error) {
data, err := json.Marshal(payload.Value)
if err != nil {
return models.RemoteSessionPayload{}, err
}
var request models.RemoteSessionPayload
if err := json.Unmarshal(data, &request); err != nil {
return models.RemoteSessionPayload{}, err
}
if request.SessionID == "" || len(request.SessionID) > 128 {
return models.RemoteSessionPayload{}, errors.New("invalid remote session id")
}
return request, nil
}
func normalizeTerminalSize(rows uint16, columns uint16) (uint16, uint16) {
if rows < 5 {
rows = 24
}
if rows > 200 {
rows = 200
}
if columns < 20 {
columns = 80
}
if columns > 400 {
columns = 400
}
return rows, columns
}
func HandleRemoteSessionOpen(client paho.Client, hubKey string, payload models.Payload, authenticated bool, configuration *models.Configuration) {
request, err := decodeRemotePayload(payload)
if err != nil {
return
}
if accessErr := validateRemoteAccess(authenticated); accessErr != nil {
publishRemoteStatus(client, hubKey, configuration, request.SessionID, request.Kind, "error", accessErr.Error())
return
}
switch request.Kind {
case "logs":
err = remoteAccess.openLogs(client, hubKey, configuration, request)
case "shell":
err = remoteAccess.openShell(client, hubKey, configuration, request)
default:
err = errors.New("unsupported remote session kind")
}
if err != nil {
if errors.Is(err, errRemoteSessionExists) {
publishRemoteStatus(client, hubKey, configuration, request.SessionID, request.Kind, "opened", "")
return
}
publishRemoteStatus(client, hubKey, configuration, request.SessionID, request.Kind, "error", err.Error())
}
}
func HandleRemoteSessionInput(client paho.Client, hubKey string, payload models.Payload, authenticated bool, configuration *models.Configuration) {
if validateRemoteAccess(authenticated) != nil {
return
}
request, err := decodeRemotePayload(payload)
if err != nil || len(request.Data) > remoteInputLimit*2 {
return
}
data, err := base64.StdEncoding.DecodeString(request.Data)
if err != nil || len(data) > remoteInputLimit {
return
}
remoteAccess.mu.Lock()
session := remoteAccess.sessions[request.SessionID]
remoteAccess.mu.Unlock()
if session == nil || session.kind != "shell" || session.pty == nil {
publishRemoteStatus(client, hubKey, configuration, request.SessionID, "shell", "error", "remote session is not open")
return
}
if _, err := session.pty.Write(data); err != nil {
publishRemoteStatus(client, hubKey, configuration, request.SessionID, "shell", "error", "failed to write terminal input")
}
}
func HandleRemoteSessionResize(payload models.Payload, authenticated bool) {
if validateRemoteAccess(authenticated) != nil {
return
}
request, err := decodeRemotePayload(payload)
if err != nil {
return
}
rows, columns := normalizeTerminalSize(request.Rows, request.Columns)
remoteAccess.mu.Lock()
session := remoteAccess.sessions[request.SessionID]
remoteAccess.mu.Unlock()
if session != nil && session.kind == "shell" && session.pty != nil {
_ = pty.Setsize(session.pty, &pty.Winsize{Rows: rows, Cols: columns})
}
}
func HandleRemoteSessionClose(payload models.Payload, authenticated bool) {
if !authenticated {
return
}
request, err := decodeRemotePayload(payload)
if err == nil {
remoteAccess.close(request.SessionID)
}
}
func (manager *remoteAccessManager) reserve(session *remoteSession) ([]string, error) {
manager.mu.Lock()
defer manager.mu.Unlock()
if _, exists := manager.sessions[session.id]; exists {
return nil, errRemoteSessionExists
}
if len(manager.sessions) >= remoteSessionLimit {
return nil, errRemoteSessionLimit
}
manager.sessions[session.id] = session
history := append([]string(nil), manager.history...)
return history, nil
}
func (manager *remoteAccessManager) expire(client paho.Client, hubKey string, configuration *models.Configuration, session *remoteSession) {
session.timer = time.AfterFunc(remoteSessionLifetime, func() {
manager.close(session.id)
if session.kind == "logs" {
publishRemoteStatus(client, hubKey, configuration, session.id, session.kind, "closed", "session lifetime reached")
}
})
}
func (manager *remoteAccessManager) openLogs(client paho.Client, hubKey string, configuration *models.Configuration, request models.RemoteSessionPayload) error {
session := &remoteSession{
id: request.SessionID,
kind: "logs",
logs: make(chan string, 256),
done: make(chan struct{}),
}
history, err := manager.reserve(session)
if err != nil {
return err
}
manager.expire(client, hubKey, configuration, session)
tail := request.Tail
if tail <= 0 || tail > remoteHistoryLimit {
tail = 200
}
if len(history) > tail {
history = history[len(history)-tail:]
}
publishRemoteStatus(client, hubKey, configuration, session.id, session.kind, "opened", "")
go func() {
for _, line := range history {
publishRemoteOutput(client, hubKey, configuration, session.id, session.kind, line)
}
for {
select {
case line := <-session.logs:
publishRemoteOutput(client, hubKey, configuration, session.id, session.kind, line)
case <-session.done:
return
}
}
}()
return nil
}
func (manager *remoteAccessManager) openShell(client paho.Client, hubKey string, configuration *models.Configuration, request models.RemoteSessionPayload) error {
rows, columns := normalizeTerminalSize(request.Rows, request.Columns)
ctx, cancel := context.WithCancel(context.Background())
session := &remoteSession{
id: request.SessionID,
kind: "shell",
cancel: cancel,
done: make(chan struct{}),
}
if _, err := manager.reserve(session); err != nil {
cancel()
return err
}
manager.expire(client, hubKey, configuration, session)
command := exec.CommandContext(ctx, "/bin/sh")
command.Env = append(os.Environ(), "TERM=xterm-256color", "HISTFILE=/dev/null")
terminal, err := pty.StartWithSize(command, &pty.Winsize{Rows: rows, Cols: columns})
if err != nil {
manager.remove(session.id)
cancel()
return err
}
session.pty = terminal
publishRemoteStatus(client, hubKey, configuration, session.id, session.kind, "opened", "")
go manager.forwardShell(client, hubKey, configuration, session, command)
return nil
}
func (manager *remoteAccessManager) forwardShell(client paho.Client, hubKey string, configuration *models.Configuration, session *remoteSession, command *exec.Cmd) {
buffer := make([]byte, remoteOutputChunkSize)
for {
count, err := session.pty.Read(buffer)
if count > 0 {
publishRemoteOutput(client, hubKey, configuration, session.id, session.kind, base64.StdEncoding.EncodeToString(buffer[:count]))
}
if err != nil {
if !errors.Is(err, io.EOF) && !errors.Is(err, os.ErrClosed) {
publishRemoteStatus(client, hubKey, configuration, session.id, session.kind, "error", "terminal stream closed unexpectedly")
}
break
}
}
_ = command.Wait()
manager.remove(session.id)
publishRemoteStatus(client, hubKey, configuration, session.id, session.kind, "closed", "")
}
func (manager *remoteAccessManager) close(sessionID string) {
manager.mu.Lock()
session := manager.sessions[sessionID]
delete(manager.sessions, sessionID)
manager.mu.Unlock()
if session == nil {
return
}
if session.cancel != nil {
session.cancel()
}
if session.pty != nil {
_ = session.pty.Close()
}
if session.logs != nil {
close(session.done)
}
if session.timer != nil {
session.timer.Stop()
}
}
func (manager *remoteAccessManager) remove(sessionID string) {
manager.mu.Lock()
session := manager.sessions[sessionID]
delete(manager.sessions, sessionID)
manager.mu.Unlock()
if session != nil && session.timer != nil {
session.timer.Stop()
}
}
func publishRemoteStatus(client paho.Client, hubKey string, configuration *models.Configuration, sessionID string, kind string, state string, errorMessage string) {
status := models.RemoteSessionStatus{
Timestamp: time.Now().Unix(),
SessionID: sessionID,
Kind: kind,
State: state,
Error: errorMessage,
}
value, _ := json.Marshal(status)
var statusValue map[string]interface{}
_ = json.Unmarshal(value, &statusValue)
publishRemote(client, hubKey, configuration, "remote-session-status", statusValue, 1)
}
func publishRemoteOutput(client paho.Client, hubKey string, configuration *models.Configuration, sessionID string, kind string, data string) {
publishRemote(client, hubKey, configuration, "remote-session-output", map[string]interface{}{
"timestamp": time.Now().Unix(),
"session_id": sessionID,
"kind": kind,
"data": data,
}, 0)
}
func publishRemote(client paho.Client, hubKey string, configuration *models.Configuration, action string, value map[string]interface{}, qos byte) {
message := models.Message{Payload: models.Payload{
Action: action,
DeviceId: configuration.Config.Key,
Value: value,
}}
payload, err := models.PackageMQTTMessage(configuration, message)
if err == nil {
client.Publish("kerberos/hub/"+hubKey, qos, false, payload)
}
}