From da91f8e592b04ec083357aa91f835830f4025ab9 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Do=C4=9Fan=20Can=20Bak=C4=B1r?= Date: Thu, 26 Feb 2026 17:02:36 +0300 Subject: [PATCH 1/2] fix: interaction data corruption causing unmarshal errors --- pkg/client/client.go | 1 + pkg/server/dns_server.go | 17 +- pkg/server/ftp_server.go | 9 +- pkg/server/http_server.go | 17 +- pkg/server/ldap_server.go | 17 +- pkg/server/responder_server.go | 9 +- pkg/server/smb_server.go | 9 +- pkg/server/smtp_server.go | 17 +- pkg/storage/roundtrip_test.go | 444 +++++++++++++++++++++++++++++++++ pkg/storage/storagedb.go | 10 +- 10 files changed, 497 insertions(+), 53 deletions(-) create mode 100644 pkg/storage/roundtrip_test.go diff --git a/pkg/client/client.go b/pkg/client/client.go index 929a8cad..a45fd75e 100644 --- a/pkg/client/client.go +++ b/pkg/client/client.go @@ -459,6 +459,7 @@ func (c *Client) getInteractions(callback InteractionCallback) error { gologger.Error().Msgf("Could not decrypt interaction: %v\n", err) continue } + plaintext = bytes.TrimRight(plaintext, " \t\r\n") interaction := &server.Interaction{} if err := jsoniter.Unmarshal(plaintext, interaction); err != nil { gologger.Error().Msgf("Could not unmarshal interaction data interaction: %v\n", err) diff --git a/pkg/server/dns_server.go b/pkg/server/dns_server.go index 21388c67..b055cc5f 100644 --- a/pkg/server/dns_server.go +++ b/pkg/server/dns_server.go @@ -1,7 +1,6 @@ package server import ( - "bytes" "context" "fmt" "net" @@ -295,12 +294,12 @@ func (h *DNSServer) handleInteraction(domain string, w dns.ResponseWriter, r *dn h.options.OnResult(interaction) } - buffer := &bytes.Buffer{} - if err := jsoniter.NewEncoder(buffer).Encode(interaction); err != nil { + data, err := jsoniter.Marshal(interaction) + if err != nil { gologger.Warning().Msgf("Could not encode root tld dns interaction: %s\n", err) } else { - gologger.Debug().Msgf("Root TLD DNS Interaction: \n%s\n", buffer.String()) - if err := h.options.Storage.AddInteractionWithId(correlationID, buffer.Bytes()); err != nil { + gologger.Debug().Msgf("Root TLD DNS Interaction: \n%s\n", string(data)) + if err := h.options.Storage.AddInteractionWithId(correlationID, data); err != nil { gologger.Warning().Msgf("Could not store dns interaction: %s\n", err) } } @@ -348,12 +347,12 @@ func (h *DNSServer) handleInteraction(domain string, w dns.ResponseWriter, r *dn RemoteAddress: host, Timestamp: time.Now(), } - buffer := &bytes.Buffer{} - if err := jsoniter.NewEncoder(buffer).Encode(interaction); err != nil { + data, err := jsoniter.Marshal(interaction) + if err != nil { gologger.Warning().Msgf("Could not encode dns interaction: %s\n", err) } else { - gologger.Debug().Msgf("DNS Interaction: \n%s\n", buffer.String()) - if err := h.options.Storage.AddInteraction(correlationID, buffer.Bytes()); err != nil { + gologger.Debug().Msgf("DNS Interaction: \n%s\n", string(data)) + if err := h.options.Storage.AddInteraction(correlationID, data); err != nil { gologger.Warning().Msgf("Could not store dns interaction: %s\n", err) } } diff --git a/pkg/server/ftp_server.go b/pkg/server/ftp_server.go index eeb60b13..e43f466f 100644 --- a/pkg/server/ftp_server.go +++ b/pkg/server/ftp_server.go @@ -1,7 +1,6 @@ package server import ( - "bytes" "crypto/tls" "fmt" "io" @@ -131,12 +130,12 @@ func (h *FTPServer) recordInteraction(remoteAddress, data string) { RawRequest: data, Timestamp: time.Now(), } - buffer := &bytes.Buffer{} - if err := jsoniter.NewEncoder(buffer).Encode(interaction); err != nil { + dataBytes, err := jsoniter.Marshal(interaction) + if err != nil { gologger.Warning().Msgf("Could not encode ftp interaction: %s\n", err) } else { - gologger.Debug().Msgf("FTP Interaction: \n%s\n", buffer.String()) - if err := h.options.Storage.AddInteractionWithId(h.options.Token, buffer.Bytes()); err != nil { + gologger.Debug().Msgf("FTP Interaction: \n%s\n", string(dataBytes)) + if err := h.options.Storage.AddInteractionWithId(h.options.Token, dataBytes); err != nil { gologger.Warning().Msgf("Could not store ftp interaction: %s\n", err) } } diff --git a/pkg/server/http_server.go b/pkg/server/http_server.go index a0a06d1e..b5b7138b 100644 --- a/pkg/server/http_server.go +++ b/pkg/server/http_server.go @@ -1,7 +1,6 @@ package server import ( - "bytes" "crypto/tls" "encoding/base64" "fmt" @@ -158,12 +157,12 @@ func (h *HTTPServer) logger(handler http.Handler) http.HandlerFunc { RemoteAddress: host, Timestamp: time.Now(), } - buffer := &bytes.Buffer{} - if err := jsoniter.NewEncoder(buffer).Encode(interaction); err != nil { + data, err := jsoniter.Marshal(interaction) + if err != nil { gologger.Warning().Msgf("Could not encode root tld http interaction: %s\n", err) } else { - gologger.Debug().Msgf("Root TLD HTTP Interaction: \n%s\n", buffer.String()) - if err := h.options.Storage.AddInteractionWithId(ID, buffer.Bytes()); err != nil { + gologger.Debug().Msgf("Root TLD HTTP Interaction: \n%s\n", string(data)) + if err := h.options.Storage.AddInteractionWithId(ID, data); err != nil { gologger.Warning().Msgf("Could not store root tld http interaction: %s\n", err) } } @@ -218,13 +217,13 @@ func (h *HTTPServer) handleInteraction(r *http.Request, uniqueID, fullID, reqStr RemoteAddress: hostPort, Timestamp: time.Now(), } - buffer := &bytes.Buffer{} - if err := jsoniter.NewEncoder(buffer).Encode(interaction); err != nil { + data, err := jsoniter.Marshal(interaction) + if err != nil { gologger.Warning().Msgf("Could not encode http interaction: %s\n", err) } else { - gologger.Debug().Msgf("HTTP Interaction: \n%s\n", buffer.String()) + gologger.Debug().Msgf("HTTP Interaction: \n%s\n", string(data)) - if err := h.options.Storage.AddInteraction(correlationID, buffer.Bytes()); err != nil { + if err := h.options.Storage.AddInteraction(correlationID, data); err != nil { gologger.Warning().Msgf("Could not store http interaction: %s\n", err) } } diff --git a/pkg/server/ldap_server.go b/pkg/server/ldap_server.go index bf60add5..35d430b5 100644 --- a/pkg/server/ldap_server.go +++ b/pkg/server/ldap_server.go @@ -1,7 +1,6 @@ package server import ( - "bytes" "crypto/tls" "fmt" "strings" @@ -145,12 +144,12 @@ func (ldapServer *LDAPServer) handleInteraction(uniqueID, fullID, reqString, hos RemoteAddress: host, Timestamp: time.Now(), } - buffer := &bytes.Buffer{} - if err := jsoniter.NewEncoder(buffer).Encode(interaction); err != nil { + data, err := jsoniter.Marshal(interaction) + if err != nil { gologger.Warning().Msgf("Could not encode ldap interaction: %s\n", err) } else { - gologger.Debug().Msgf("LDAP Interaction: \n%s\n", buffer.String()) - if err := ldapServer.options.Storage.AddInteraction(correlationID, buffer.Bytes()); err != nil { + gologger.Debug().Msgf("LDAP Interaction: \n%s\n", string(data)) + if err := ldapServer.options.Storage.AddInteraction(correlationID, data); err != nil { gologger.Warning().Msgf("Could not store ldap interaction: %s\n", err) } } @@ -411,12 +410,12 @@ func (ldapServer *LDAPServer) logInteraction(interaction Interaction) { // Correlation id doesn't apply here, we skip encryption interaction.Protocol = "ldap" interaction.Timestamp = time.Now() - buffer := &bytes.Buffer{} - if err := jsoniter.NewEncoder(buffer).Encode(interaction); err != nil { + data, err := jsoniter.Marshal(interaction) + if err != nil { gologger.Warning().Msgf("Could not encode ldap interaction: %s\n", err) } else { - gologger.Debug().Msgf("LDAP Interaction: \n%s\n", buffer.String()) - if err := ldapServer.options.Storage.AddInteractionWithId(ldapServer.options.Token, buffer.Bytes()); err != nil { + gologger.Debug().Msgf("LDAP Interaction: \n%s\n", string(data)) + if err := ldapServer.options.Storage.AddInteractionWithId(ldapServer.options.Token, data); err != nil { gologger.Warning().Msgf("Could not store ldap interaction: %s\n", err) } } diff --git a/pkg/server/responder_server.go b/pkg/server/responder_server.go index f13fa88e..20ff92c4 100644 --- a/pkg/server/responder_server.go +++ b/pkg/server/responder_server.go @@ -1,7 +1,6 @@ package server import ( - "bytes" "os" "os/exec" "path/filepath" @@ -89,12 +88,12 @@ func (h *ResponderServer) ListenAndServe(responderAlive chan bool) error { RawRequest: responderData, Timestamp: time.Now(), } - buffer := &bytes.Buffer{} - if err := jsoniter.NewEncoder(buffer).Encode(interaction); err != nil { + data, err := jsoniter.Marshal(interaction) + if err != nil { gologger.Warning().Msgf("Could not encode responder interaction: %s\n", err) } else { - gologger.Debug().Msgf("Responder Interaction: \n%s\n", buffer.String()) - if err := h.options.Storage.AddInteractionWithId(h.options.Token, buffer.Bytes()); err != nil { + gologger.Debug().Msgf("Responder Interaction: \n%s\n", string(data)) + if err := h.options.Storage.AddInteractionWithId(h.options.Token, data); err != nil { gologger.Warning().Msgf("Could not store dns interaction: %s\n", err) } } diff --git a/pkg/server/smb_server.go b/pkg/server/smb_server.go index a6b7ef83..ac445352 100644 --- a/pkg/server/smb_server.go +++ b/pkg/server/smb_server.go @@ -1,7 +1,6 @@ package server import ( - "bytes" "fmt" "os" "os/exec" @@ -101,12 +100,12 @@ func (h *SMBServer) ListenAndServe(smbAlive chan bool) error { RawRequest: smbData, Timestamp: time.Now(), } - buffer := &bytes.Buffer{} - if err := jsoniter.NewEncoder(buffer).Encode(interaction); err != nil { + data, err := jsoniter.Marshal(interaction) + if err != nil { gologger.Warning().Msgf("Could not encode smb interaction: %s\n", err) } else { - gologger.Debug().Msgf("SMB Interaction: \n%s\n", buffer.String()) - if err := h.options.Storage.AddInteractionWithId(h.options.Token, buffer.Bytes()); err != nil { + gologger.Debug().Msgf("SMB Interaction: \n%s\n", string(data)) + if err := h.options.Storage.AddInteractionWithId(h.options.Token, data); err != nil { gologger.Warning().Msgf("Could not store dns interaction: %s\n", err) } } diff --git a/pkg/server/smtp_server.go b/pkg/server/smtp_server.go index 4e981ef1..6c88ca06 100644 --- a/pkg/server/smtp_server.go +++ b/pkg/server/smtp_server.go @@ -1,7 +1,6 @@ package server import ( - "bytes" "crypto/tls" "net" "strings" @@ -107,12 +106,12 @@ func (h *SMTPServer) defaultHandler(remoteAddr net.Addr, from string, to []strin RemoteAddress: host, Timestamp: time.Now(), } - buffer := &bytes.Buffer{} - if err := jsoniter.NewEncoder(buffer).Encode(interaction); err != nil { + data, err := jsoniter.Marshal(interaction) + if err != nil { gologger.Warning().Msgf("Could not encode root tld SMTP interaction: %s\n", err) } else { - gologger.Debug().Msgf("Root TLD SMTP Interaction: \n%s\n", buffer.String()) - if err := h.options.Storage.AddInteractionWithId(ID, buffer.Bytes()); err != nil { + gologger.Debug().Msgf("Root TLD SMTP Interaction: \n%s\n", string(data)) + if err := h.options.Storage.AddInteractionWithId(ID, data); err != nil { gologger.Warning().Msgf("Could not store root tld smtp interaction: %s\n", err) } } @@ -148,12 +147,12 @@ func (h *SMTPServer) defaultHandler(remoteAddr net.Addr, from string, to []strin RemoteAddress: host, Timestamp: time.Now(), } - buffer := &bytes.Buffer{} - if err := jsoniter.NewEncoder(buffer).Encode(interaction); err != nil { + data, err := jsoniter.Marshal(interaction) + if err != nil { gologger.Warning().Msgf("Could not encode smtp interaction: %s\n", err) } else { - gologger.Debug().Msgf("%s\n", buffer.String()) - if err := h.options.Storage.AddInteraction(correlationID, buffer.Bytes()); err != nil { + gologger.Debug().Msgf("%s\n", string(data)) + if err := h.options.Storage.AddInteraction(correlationID, data); err != nil { gologger.Warning().Msgf("Could not store smtp interaction: %s\n", err) } } diff --git a/pkg/storage/roundtrip_test.go b/pkg/storage/roundtrip_test.go new file mode 100644 index 00000000..7113871f --- /dev/null +++ b/pkg/storage/roundtrip_test.go @@ -0,0 +1,444 @@ +package storage + +import ( + "bytes" + "crypto/aes" + "crypto/cipher" + "crypto/rand" + "crypto/rsa" + "crypto/sha256" + "crypto/x509" + "encoding/base64" + "encoding/pem" + "os" + "testing" + "time" + + jsoniter "github.com/json-iterator/go" + "github.com/google/uuid" + "github.com/rs/xid" + "github.com/stretchr/testify/require" +) + +// interaction mirrors server.Interaction to avoid circular import +type interaction struct { + Protocol string `json:"protocol"` + UniqueID string `json:"unique-id"` + FullId string `json:"full-id"` + QType string `json:"q-type,omitempty"` + RawRequest string `json:"raw-request,omitempty"` + RawResponse string `json:"raw-response,omitempty"` + SMTPFrom string `json:"smtp-from,omitempty"` + RemoteAddress string `json:"remote-address"` + Timestamp time.Time `json:"timestamp"` + AsnInfo []map[string]string `json:"asninfo,omitempty"` +} + +// Realistic DNS message dump (from miekg/dns String() output) +const dnsRequest = `;; opcode: QUERY, status: NOERROR, id: 12345 +;; flags: rd; QUERY: 1, ANSWER: 0, AUTHORITY: 0, ADDITIONAL: 1 + +;; OPT PSEUDOSECTION: +; EDNS: version 0; flags: ; udp: 4096 + +;; QUESTION SECTION: +;abc123def456ghi.oast.fun. IN A +` + +const dnsResponse = `;; opcode: QUERY, status: NOERROR, id: 12345 +;; flags: qr aa rd; QUERY: 1, ANSWER: 1, AUTHORITY: 2, ADDITIONAL: 2 + +;; QUESTION SECTION: +;abc123def456ghi.oast.fun. IN A + +;; ANSWER SECTION: +abc123def456ghi.oast.fun. 3600 IN A 1.2.3.4 + +;; AUTHORITY SECTION: +oast.fun. 3600 IN NS ns1.oast.fun. +oast.fun. 3600 IN NS ns2.oast.fun. + +;; ADDITIONAL SECTION: +ns1.oast.fun. 3600 IN A 1.2.3.4 +ns2.oast.fun. 3600 IN A 1.2.3.4 +` + +func generateRSAKeyPair(t *testing.T) (*rsa.PrivateKey, string) { + t.Helper() + priv, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + pubBytes, err := x509.MarshalPKIXPublicKey(&priv.PublicKey) + require.NoError(t, err) + pubPem := pem.EncodeToMemory(&pem.Block{Type: "RSA PUBLIC KEY", Bytes: pubBytes}) + return priv, base64.StdEncoding.EncodeToString(pubPem) +} + +func clientDecrypt(t *testing.T, priv *rsa.PrivateKey, aesKeyEncrypted string, cipherData string) []byte { + t.Helper() + decodedKey, err := base64.StdEncoding.DecodeString(aesKeyEncrypted) + require.NoError(t, err) + keyPlaintext, err := rsa.DecryptOAEP(sha256.New(), rand.Reader, priv, decodedKey, nil) + require.NoError(t, err) + cipherText, err := base64.StdEncoding.DecodeString(cipherData) + require.NoError(t, err) + require.Greater(t, len(cipherText), aes.BlockSize) + iv := cipherText[:aes.BlockSize] + cipherText = cipherText[aes.BlockSize:] + block, err := aes.NewCipher(keyPlaintext) + require.NoError(t, err) + stream := cipher.NewCTR(block, iv) + decoded := make([]byte, len(cipherText)) + stream.XORKeyStream(decoded, cipherText) + return decoded +} + +func TestFullRoundTripInMemory(t *testing.T) { + mem, err := New(&Options{EvictionTTL: 1 * time.Hour}) + require.NoError(t, err) + defer mem.Close() + + priv, pubKeyB64 := generateRSAKeyPair(t) + secret := uuid.New().String() + correlationID := xid.New().String() + + err = mem.SetIDPublicKey(correlationID, secret, pubKeyB64) + require.NoError(t, err) + + // Create and store 3 DNS interactions (like dns_server.go does) + for i := 0; i < 3; i++ { + inter := &interaction{ + Protocol: "dns", + UniqueID: "abc123def456ghi", + FullId: "abc123def456ghi.oast.fun", + QType: "A", + RawRequest: dnsRequest, + RawResponse: dnsResponse, + RemoteAddress: "10.0.0.1", + Timestamp: time.Now(), + } + buffer := &bytes.Buffer{} + err = jsoniter.NewEncoder(buffer).Encode(inter) + require.NoError(t, err, "encode interaction %d", i) + + err = mem.AddInteraction(correlationID, buffer.Bytes()) + require.NoError(t, err, "add interaction %d", i) + } + + // Retrieve + data, aesKey, err := mem.GetInteractions(correlationID, secret) + require.NoError(t, err) + require.Len(t, data, 3) + + // Decrypt and unmarshal each (like client.go does) + for i, d := range data { + plaintext := clientDecrypt(t, priv, aesKey, d) + result := &interaction{} + err = jsoniter.Unmarshal(plaintext, result) + require.NoError(t, err, "unmarshal interaction %d: plaintext=%q", i, string(plaintext[:min(len(plaintext), 100)])) + require.Equal(t, "dns", result.Protocol) + require.Equal(t, "abc123def456ghi", result.UniqueID) + require.Equal(t, dnsRequest, result.RawRequest) + require.Equal(t, dnsResponse, result.RawResponse) + } +} + +func TestFullRoundTripDisk(t *testing.T) { + tmpDir, err := os.MkdirTemp("", "interactsh-test-*") + require.NoError(t, err) + defer os.RemoveAll(tmpDir) + + db, err := New(&Options{EvictionTTL: 1 * time.Hour, DbPath: tmpDir}) + require.NoError(t, err) + defer db.Close() + + priv, pubKeyB64 := generateRSAKeyPair(t) + secret := uuid.New().String() + correlationID := xid.New().String() + + err = db.SetIDPublicKey(correlationID, secret, pubKeyB64) + require.NoError(t, err) + + // Create and store 3 DNS interactions + for i := 0; i < 3; i++ { + inter := &interaction{ + Protocol: "dns", + UniqueID: "abc123def456ghi", + FullId: "abc123def456ghi.oast.fun", + QType: "A", + RawRequest: dnsRequest, + RawResponse: dnsResponse, + RemoteAddress: "10.0.0.1", + Timestamp: time.Now(), + } + buffer := &bytes.Buffer{} + err = jsoniter.NewEncoder(buffer).Encode(inter) + require.NoError(t, err, "encode interaction %d", i) + + err = db.AddInteraction(correlationID, buffer.Bytes()) + require.NoError(t, err, "add interaction %d", i) + } + + // Retrieve + data, aesKey, err := db.GetInteractions(correlationID, secret) + require.NoError(t, err) + require.Len(t, data, 3) + + // Decrypt and unmarshal each + for i, d := range data { + plaintext := clientDecrypt(t, priv, aesKey, d) + result := &interaction{} + err = jsoniter.Unmarshal(plaintext, result) + require.NoError(t, err, "unmarshal interaction %d: plaintext=%q", i, string(plaintext[:min(len(plaintext), 100)])) + require.Equal(t, "dns", result.Protocol) + require.Equal(t, "abc123def456ghi", result.UniqueID) + require.Equal(t, dnsRequest, result.RawRequest) + require.Equal(t, dnsResponse, result.RawResponse) + } +} + +// TestPollResponseRoundTrip tests the full flow including the HTTP PollResponse JSON encoding +func TestPollResponseRoundTrip(t *testing.T) { + type PollResponse struct { + Data []string `json:"data"` + Extra []string `json:"extra"` + AESKey string `json:"aes_key"` + TLDData []string `json:"tlddata,omitempty"` + } + + mem, err := New(&Options{EvictionTTL: 1 * time.Hour}) + require.NoError(t, err) + defer mem.Close() + + priv, pubKeyB64 := generateRSAKeyPair(t) + secret := uuid.New().String() + correlationID := xid.New().String() + + err = mem.SetIDPublicKey(correlationID, secret, pubKeyB64) + require.NoError(t, err) + + // Store a DNS interaction + inter := &interaction{ + Protocol: "dns", + UniqueID: "abc123def456ghi", + FullId: "abc123def456ghi.oast.fun", + QType: "A", + RawRequest: dnsRequest, + RawResponse: dnsResponse, + RemoteAddress: "10.0.0.1", + Timestamp: time.Now(), + } + buffer := &bytes.Buffer{} + err = jsoniter.NewEncoder(buffer).Encode(inter) + require.NoError(t, err) + err = mem.AddInteraction(correlationID, buffer.Bytes()) + require.NoError(t, err) + + // Retrieve (server side) + data, aesKey, err := mem.GetInteractions(correlationID, secret) + require.NoError(t, err) + + // Simulate PollResponse JSON encoding/decoding (server→client HTTP) + response := &PollResponse{Data: data, AESKey: aesKey} + var responseBuf bytes.Buffer + err = jsoniter.NewEncoder(&responseBuf).Encode(response) + require.NoError(t, err) + + // Decode on client side + receivedResponse := &PollResponse{} + err = jsoniter.NewDecoder(&responseBuf).Decode(receivedResponse) + require.NoError(t, err) + + require.Len(t, receivedResponse.Data, 1) + + // Decrypt and unmarshal + plaintext := clientDecrypt(t, priv, receivedResponse.AESKey, receivedResponse.Data[0]) + result := &interaction{} + err = jsoniter.Unmarshal(plaintext, result) + require.NoError(t, err, "unmarshal failed: plaintext[:100]=%q", string(plaintext[:min(len(plaintext), 100)])) + require.Equal(t, "dns", result.Protocol) +} + +// TestTrailingNewlineHandling verifies that the trailing \n from Encode() doesn't cause issues +func TestTrailingNewlineHandling(t *testing.T) { + inter := &interaction{ + Protocol: "dns", + UniqueID: "test", + FullId: "test.example.com", + RemoteAddress: "1.2.3.4", + Timestamp: time.Now(), + } + buffer := &bytes.Buffer{} + err := jsoniter.NewEncoder(buffer).Encode(inter) + require.NoError(t, err) + + encoded := buffer.Bytes() + t.Logf("Encoded JSON length: %d", len(encoded)) + t.Logf("Last byte: 0x%02x (newline=%v)", encoded[len(encoded)-1], encoded[len(encoded)-1] == '\n') + + // Verify trailing newline is present + require.Equal(t, byte('\n'), encoded[len(encoded)-1], "Encode() should append trailing newline") + + // Verify jsoniter.Unmarshal handles trailing newline + result := &interaction{} + err = jsoniter.Unmarshal(encoded, result) + require.NoError(t, err, "Unmarshal should handle trailing newline") + require.Equal(t, "dns", result.Protocol) +} + +// TestJsoniterControlCharacterEscaping verifies jsoniter properly escapes control characters +func TestJsoniterControlCharacterEscaping(t *testing.T) { + // DNS message String() output contains tabs and newlines + inter := &interaction{ + Protocol: "dns", + UniqueID: "test", + FullId: "test.example.com", + RawRequest: "line1\nline2\ttab\rcarriage", + RemoteAddress: "1.2.3.4", + Timestamp: time.Now(), + } + buffer := &bytes.Buffer{} + err := jsoniter.NewEncoder(buffer).Encode(inter) + require.NoError(t, err) + + encoded := buffer.Bytes() + // Check that no raw control characters (except the trailing \n) exist in the JSON + for i, b := range encoded[:len(encoded)-1] { // skip trailing \n + if b < 0x20 { + t.Errorf("Found unescaped control character 0x%02x at position %d in JSON: ...%q...", + b, i, string(encoded[max(0, i-20):min(len(encoded), i+20)])) + } + } + + // Verify round-trip + result := &interaction{} + err = jsoniter.Unmarshal(encoded, result) + require.NoError(t, err) + require.Equal(t, "line1\nline2\ttab\rcarriage", result.RawRequest) +} + +// TestStaleDataCleanupOnReRegistration verifies that stale LevelDB data from a +// previous registration (encrypted with an old AES key) is purged when the same +// correlation ID is re-registered after cache eviction. +func TestStaleDataCleanupOnReRegistration(t *testing.T) { + tmpDir, err := os.MkdirTemp("", "interactsh-stale-*") + require.NoError(t, err) + defer os.RemoveAll(tmpDir) + + // Use a very short eviction TTL so the cache entry gets evicted quickly + db, err := New(&Options{EvictionTTL: 100 * time.Millisecond, DbPath: tmpDir}) + require.NoError(t, err) + defer db.Close() + + secret := uuid.New().String() + correlationID := xid.New().String() + + // First registration + priv1, pubKey1 := generateRSAKeyPair(t) + err = db.SetIDPublicKey(correlationID, secret, pubKey1) + require.NoError(t, err) + + // Store an interaction encrypted with AES key K1 + inter := &interaction{ + Protocol: "dns", + UniqueID: "first-registration", + FullId: "first.example.com", + RemoteAddress: "1.2.3.4", + Timestamp: time.Now(), + } + data1, err := jsoniter.Marshal(inter) + require.NoError(t, err) + err = db.AddInteraction(correlationID, data1) + require.NoError(t, err) + + // Wait for cache eviction + time.Sleep(200 * time.Millisecond) + + // Verify cache entry is evicted + _, found := db.cache.GetIfPresent(correlationID) + require.False(t, found, "cache entry should be evicted") + + // Re-register with the same correlation ID (simulates session restore) + // This generates a NEW AES key K2 + priv2, pubKey2 := generateRSAKeyPair(t) + err = db.SetIDPublicKey(correlationID, secret, pubKey2) + require.NoError(t, err) + + // Store a new interaction encrypted with K2 + inter2 := &interaction{ + Protocol: "dns", + UniqueID: "second-registration", + FullId: "second.example.com", + RemoteAddress: "5.6.7.8", + Timestamp: time.Now(), + } + data2, err := jsoniter.Marshal(inter2) + require.NoError(t, err) + err = db.AddInteraction(correlationID, data2) + require.NoError(t, err) + + // Retrieve interactions + interactions, aesKey, err := db.GetInteractions(correlationID, secret) + require.NoError(t, err) + + // Should only have 1 interaction (the new one), not the stale one + require.Len(t, interactions, 1, "stale data from first registration should have been purged") + + // Decrypt and verify it's the second interaction + plaintext := clientDecrypt(t, priv2, aesKey, interactions[0]) + result := &interaction{} + err = jsoniter.Unmarshal(plaintext, result) + require.NoError(t, err, "should unmarshal successfully with new key") + require.Equal(t, "second-registration", result.UniqueID) + + // priv1 is no longer useful (old key pair) + _ = priv1 +} + +// TestCacheEvictionCleansLevelDB verifies the OnCacheRemovalCallback properly +// deletes LevelDB entries when cache entries are evicted. +func TestCacheEvictionCleansLevelDB(t *testing.T) { + tmpDir, err := os.MkdirTemp("", "interactsh-eviction-*") + require.NoError(t, err) + defer os.RemoveAll(tmpDir) + + db, err := New(&Options{EvictionTTL: 100 * time.Millisecond, DbPath: tmpDir}) + require.NoError(t, err) + defer db.Close() + + secret := uuid.New().String() + correlationID := xid.New().String() + + _, pubKey := generateRSAKeyPair(t) + err = db.SetIDPublicKey(correlationID, secret, pubKey) + require.NoError(t, err) + + // Store an interaction + inter := &interaction{ + Protocol: "dns", + UniqueID: "test", + FullId: "test.example.com", + RemoteAddress: "1.2.3.4", + Timestamp: time.Now(), + } + data, err := jsoniter.Marshal(inter) + require.NoError(t, err) + err = db.AddInteraction(correlationID, data) + require.NoError(t, err) + + // Verify data exists in LevelDB + raw, err := db.db.Get([]byte(correlationID), nil) + require.NoError(t, err) + require.NotEmpty(t, raw) + + // Wait for cache eviction + time.Sleep(200 * time.Millisecond) + + // Force cache cleanup (GetIfPresent triggers lazy eviction) + db.cache.GetIfPresent(correlationID) + // Small delay for async eviction callback + time.Sleep(50 * time.Millisecond) + + // LevelDB entry should be cleaned up by OnCacheRemovalCallback + _, err = db.db.Get([]byte(correlationID), nil) + require.Error(t, err, "LevelDB entry should be deleted after cache eviction") +} diff --git a/pkg/storage/storagedb.go b/pkg/storage/storagedb.go index 019f9822..b235da90 100644 --- a/pkg/storage/storagedb.go +++ b/pkg/storage/storagedb.go @@ -74,8 +74,8 @@ func New(options *Options) (*StorageDB, error) { } func (s *StorageDB) OnCacheRemovalCallback(key cache.Key, value cache.Value) { - if key, ok := value.([]byte); ok { - _ = s.db.Delete(key, &opt.WriteOptions{}) + if k, ok := key.(string); ok { + _ = s.db.Delete([]byte(k), &opt.WriteOptions{}) } } @@ -121,6 +121,12 @@ func (s *StorageDB) SetIDPublicKey(correlationID, secretKey, publicKey string) e AESKey: aesKey, AESKeyEncrypted: base64.StdEncoding.EncodeToString(ciphertext), } + // Clear any stale data from a previous registration (e.g. after cache eviction + // and session restore). Old data would be encrypted with a different AES key + // and cause decryption failures on the client. + if s.Options.UseDisk() { + _ = s.db.Delete([]byte(correlationID), nil) + } s.cache.Put(correlationID, data) return nil } From 70580fa6f3c8f332d2f50a6511885b826dad6715 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Do=C4=9Fan=20Can=20Bak=C4=B1r?= Date: Wed, 4 Mar 2026 16:16:27 +0300 Subject: [PATCH 2/2] use jsoniter.Marshal in round-trip tests --- pkg/storage/roundtrip_test.go | 15 ++++++--------- 1 file changed, 6 insertions(+), 9 deletions(-) diff --git a/pkg/storage/roundtrip_test.go b/pkg/storage/roundtrip_test.go index 7113871f..557a155b 100644 --- a/pkg/storage/roundtrip_test.go +++ b/pkg/storage/roundtrip_test.go @@ -116,11 +116,10 @@ func TestFullRoundTripInMemory(t *testing.T) { RemoteAddress: "10.0.0.1", Timestamp: time.Now(), } - buffer := &bytes.Buffer{} - err = jsoniter.NewEncoder(buffer).Encode(inter) + data, err := jsoniter.Marshal(inter) require.NoError(t, err, "encode interaction %d", i) - err = mem.AddInteraction(correlationID, buffer.Bytes()) + err = mem.AddInteraction(correlationID, data) require.NoError(t, err, "add interaction %d", i) } @@ -170,11 +169,10 @@ func TestFullRoundTripDisk(t *testing.T) { RemoteAddress: "10.0.0.1", Timestamp: time.Now(), } - buffer := &bytes.Buffer{} - err = jsoniter.NewEncoder(buffer).Encode(inter) + data, err := jsoniter.Marshal(inter) require.NoError(t, err, "encode interaction %d", i) - err = db.AddInteraction(correlationID, buffer.Bytes()) + err = db.AddInteraction(correlationID, data) require.NoError(t, err, "add interaction %d", i) } @@ -227,10 +225,9 @@ func TestPollResponseRoundTrip(t *testing.T) { RemoteAddress: "10.0.0.1", Timestamp: time.Now(), } - buffer := &bytes.Buffer{} - err = jsoniter.NewEncoder(buffer).Encode(inter) + interData, err := jsoniter.Marshal(inter) require.NoError(t, err) - err = mem.AddInteraction(correlationID, buffer.Bytes()) + err = mem.AddInteraction(correlationID, interData) require.NoError(t, err) // Retrieve (server side)