feat(canary): mysql server + connection handler
Phase 8 second commit. TCP listener orchestration + per-connection
business logic.
server.go (~70 LOC):
- mysql.Run(ctx, addr, ConnectionHandler) error
Listens via net.ListenConfig (context-aware), accepts in a loop,
dispatches each connection to handler.HandleConnection in its own
goroutine, waits on context cancellation to close the listener
and drain in-flight connections via WaitGroup
- ConnectionHandler interface (single HandleConnection method) so
tests can substitute a stub without standing up real handler deps
- Logs listener bind + accept errors via slog; suppresses net.ErrClosed
on graceful shutdown so the shutdown path is silent
handler.go (~170 LOC):
- TokenLookup interface (GetByID) and EventRecorder interface (Record)
decouple the handler from Phase 9's token.Service and Phase 10's
event.Service. Production wire-up in Phase 9 (deferred — see Phase
8 rollup notes on task 8.5 deferral)
- HandleConnection writes HandshakeV10 → reads HandshakeResponse41 →
extracts username → strips canary_ prefix → looks up token →
records event with mysql_username + mysql_client_capabilities +
mysql_client_charset in event.Extra (mirrors kubeconfig's
kubectl_* forensic capture) → writes ERR_Packet 1045 → defers
close. 10-second connection deadline.
- Defense-in-depth: usernames without canary_ prefix dropped
silently, unknown tokens dropped silently, lookup errors dropped
silently — no behavior signal to a probing attacker
- Nil EventRecorder tolerated (writes ERR but skips event) — Phase 9
can wire the handler before event.Service exists if needed
Tests (~14 cases):
- server: nil-handler rejected, basic dispatch, 10 concurrent
connections, context cancellation closes listener cleanly,
invalid address rejected
- handler: handshake sent first, non-canary username drops silently,
known token records event AND sends ERR, event.Extra carries
all three kubectl-equivalent fields with correct hex formatting,
unknown token records nothing, lookup error silent, nil
EventRecorder tolerated (still sends ERR), bad auth packet
silent
- Uses real net.Listener + net.Conn pairs (not mocks) for
behavioral fidelity — TCP wire-protocol tests should hit real
sockets
go test -race -timeout=60s passes. golangci-lint clean.
This commit is contained in:
parent
712202abc0
commit
4d4d4951ba
|
|
@ -0,0 +1,186 @@
|
||||||
|
// ©AngelaMos | 2026
|
||||||
|
// handler.go
|
||||||
|
|
||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
|
"net"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/event"
|
||||||
|
"github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/token"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
connectionDeadline = 10 * time.Second
|
||||||
|
|
||||||
|
mysqlUsernamePrefix = "canary_"
|
||||||
|
|
||||||
|
extraMySQLUsername = "mysql_username"
|
||||||
|
extraMySQLCapabilities = "mysql_client_capabilities"
|
||||||
|
extraMySQLCharset = "mysql_client_charset"
|
||||||
|
|
||||||
|
capabilitiesFormat = "0x%08x"
|
||||||
|
)
|
||||||
|
|
||||||
|
type TokenLookup interface {
|
||||||
|
GetByID(ctx context.Context, id string) (*token.Token, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
type EventRecorder interface {
|
||||||
|
Record(ctx context.Context, t *token.Token, evt *event.Event) error
|
||||||
|
}
|
||||||
|
|
||||||
|
type Handler struct {
|
||||||
|
tokens TokenLookup
|
||||||
|
events EventRecorder
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewHandler(tokens TokenLookup, events EventRecorder) *Handler {
|
||||||
|
return &Handler{tokens: tokens, events: events}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) HandleConnection(ctx context.Context, conn net.Conn) {
|
||||||
|
defer func() {
|
||||||
|
if err := conn.Close(); err != nil {
|
||||||
|
slog.WarnContext(ctx, "mysql: close connection", "error", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
if err := conn.SetDeadline(time.Now().Add(connectionDeadline)); err != nil {
|
||||||
|
slog.WarnContext(ctx, "mysql: set deadline", "error", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := h.writeHandshake(conn); err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
auth, err := ReadClientAuth(conn)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if !strings.HasPrefix(auth.Username, mysqlUsernamePrefix) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
tokenID := strings.TrimPrefix(auth.Username, mysqlUsernamePrefix)
|
||||||
|
|
||||||
|
tok, err := h.tokens.GetByID(ctx, tokenID)
|
||||||
|
if err != nil || tok == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
sourceHost := remoteHost(conn)
|
||||||
|
|
||||||
|
if h.events != nil {
|
||||||
|
if recErr := h.recordEvent(ctx, tok, sourceHost, auth); recErr != nil {
|
||||||
|
slog.WarnContext(
|
||||||
|
ctx,
|
||||||
|
"mysql: record event",
|
||||||
|
"error", recErr,
|
||||||
|
"token_id", tok.ID,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if wErr := h.writeAccessDenied(
|
||||||
|
conn,
|
||||||
|
auth.Username,
|
||||||
|
sourceHost,
|
||||||
|
); wErr != nil {
|
||||||
|
slog.WarnContext(
|
||||||
|
ctx,
|
||||||
|
"mysql: write err packet",
|
||||||
|
"error", wErr,
|
||||||
|
"token_id", tok.ID,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) writeHandshake(conn net.Conn) error {
|
||||||
|
connID, err := NewRandomConnectionID()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("connection id: %w", err)
|
||||||
|
}
|
||||||
|
authData, err := NewRandomAuthData()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("auth data: %w", err)
|
||||||
|
}
|
||||||
|
pkt, err := BuildHandshakeV10(connID, authData)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("build handshake: %w", err)
|
||||||
|
}
|
||||||
|
if _, err := conn.Write(pkt); err != nil {
|
||||||
|
return fmt.Errorf("write handshake: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) writeAccessDenied(
|
||||||
|
conn net.Conn,
|
||||||
|
username, sourceHost string,
|
||||||
|
) error {
|
||||||
|
pkt, err := BuildAccessDeniedErr(username, sourceHost)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("build err packet: %w", err)
|
||||||
|
}
|
||||||
|
if _, err := conn.Write(pkt); err != nil {
|
||||||
|
return fmt.Errorf("write err packet: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *Handler) recordEvent(
|
||||||
|
ctx context.Context,
|
||||||
|
tok *token.Token,
|
||||||
|
sourceHost string,
|
||||||
|
auth *ClientAuth,
|
||||||
|
) error {
|
||||||
|
extra, err := buildMySQLExtra(auth)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("build extra: %w", err)
|
||||||
|
}
|
||||||
|
evt := &event.Event{
|
||||||
|
TokenID: tok.ID,
|
||||||
|
SourceIP: sourceHost,
|
||||||
|
Extra: extra,
|
||||||
|
}
|
||||||
|
if err := h.events.Record(ctx, tok, evt); err != nil {
|
||||||
|
return fmt.Errorf("record: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildMySQLExtra(auth *ClientAuth) (json.RawMessage, error) {
|
||||||
|
extra := map[string]any{
|
||||||
|
extraMySQLUsername: auth.Username,
|
||||||
|
extraMySQLCapabilities: fmt.Sprintf(
|
||||||
|
capabilitiesFormat,
|
||||||
|
auth.Capabilities,
|
||||||
|
),
|
||||||
|
extraMySQLCharset: auth.Charset,
|
||||||
|
}
|
||||||
|
body, err := json.Marshal(extra)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("marshal mysql extra: %w", err)
|
||||||
|
}
|
||||||
|
return body, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func remoteHost(conn net.Conn) string {
|
||||||
|
addr := conn.RemoteAddr()
|
||||||
|
if addr == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
host, _, err := net.SplitHostPort(addr.String())
|
||||||
|
if err != nil {
|
||||||
|
return addr.String()
|
||||||
|
}
|
||||||
|
return host
|
||||||
|
}
|
||||||
|
|
@ -0,0 +1,389 @@
|
||||||
|
// ©AngelaMos | 2026
|
||||||
|
// handler_test.go
|
||||||
|
|
||||||
|
package mysql_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/binary"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"net"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/event"
|
||||||
|
"github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/token"
|
||||||
|
"github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/token/generators/mysql"
|
||||||
|
)
|
||||||
|
|
||||||
|
func closeQuietly(t *testing.T, closers ...io.Closer) {
|
||||||
|
t.Helper()
|
||||||
|
for _, c := range closers {
|
||||||
|
if err := c.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
|
||||||
|
t.Logf("close: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type fakeTokenLookup struct {
|
||||||
|
tokens map[string]*token.Token
|
||||||
|
err error
|
||||||
|
calls int32
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *fakeTokenLookup) GetByID(
|
||||||
|
_ context.Context,
|
||||||
|
id string,
|
||||||
|
) (*token.Token, error) {
|
||||||
|
atomic.AddInt32(&f.calls, 1)
|
||||||
|
if f.err != nil {
|
||||||
|
return nil, f.err
|
||||||
|
}
|
||||||
|
tok, ok := f.tokens[id]
|
||||||
|
if !ok {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
return tok, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type fakeEventRecorder struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
events []*event.Event
|
||||||
|
err error
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *fakeEventRecorder) Record(
|
||||||
|
_ context.Context,
|
||||||
|
_ *token.Token,
|
||||||
|
evt *event.Event,
|
||||||
|
) error {
|
||||||
|
f.mu.Lock()
|
||||||
|
defer f.mu.Unlock()
|
||||||
|
if f.err != nil {
|
||||||
|
return f.err
|
||||||
|
}
|
||||||
|
f.events = append(f.events, evt)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *fakeEventRecorder) snapshot() []*event.Event {
|
||||||
|
f.mu.Lock()
|
||||||
|
defer f.mu.Unlock()
|
||||||
|
out := make([]*event.Event, len(f.events))
|
||||||
|
copy(out, f.events)
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func newConnPair(t *testing.T) (clientSide, serverSide net.Conn) {
|
||||||
|
t.Helper()
|
||||||
|
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer func() {
|
||||||
|
require.NoError(t, listener.Close())
|
||||||
|
}()
|
||||||
|
|
||||||
|
type accepted struct {
|
||||||
|
conn net.Conn
|
||||||
|
err error
|
||||||
|
}
|
||||||
|
ch := make(chan accepted, 1)
|
||||||
|
go func() {
|
||||||
|
c, aErr := listener.Accept()
|
||||||
|
ch <- accepted{conn: c, err: aErr}
|
||||||
|
}()
|
||||||
|
|
||||||
|
dialed, err := net.Dial("tcp", listener.Addr().String())
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
a := <-ch
|
||||||
|
require.NoError(t, a.err)
|
||||||
|
return dialed, a.conn
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildHandshakeResponse41ForHandler(t *testing.T, username string) []byte {
|
||||||
|
t.Helper()
|
||||||
|
var payload bytes.Buffer
|
||||||
|
var caps [4]byte
|
||||||
|
binary.LittleEndian.PutUint32(caps[:], 0x0001a285)
|
||||||
|
payload.Write(caps[:])
|
||||||
|
var maxPkt [4]byte
|
||||||
|
binary.LittleEndian.PutUint32(maxPkt[:], 0x01000000)
|
||||||
|
payload.Write(maxPkt[:])
|
||||||
|
payload.WriteByte(0x21)
|
||||||
|
var filler [23]byte
|
||||||
|
payload.Write(filler[:])
|
||||||
|
payload.WriteString(username)
|
||||||
|
payload.WriteByte(0x00)
|
||||||
|
body := payload.Bytes()
|
||||||
|
out := make([]byte, 4+len(body))
|
||||||
|
n := len(body)
|
||||||
|
out[0] = byte(n & 0xff)
|
||||||
|
out[1] = byte((n >> 8) & 0xff)
|
||||||
|
out[2] = byte((n >> 16) & 0xff)
|
||||||
|
out[3] = 0x01
|
||||||
|
copy(out[4:], body)
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func readAllAvailable(t *testing.T, conn net.Conn, want int) []byte {
|
||||||
|
t.Helper()
|
||||||
|
require.NoError(t, conn.SetReadDeadline(time.Now().Add(2*time.Second)))
|
||||||
|
buf := make([]byte, want)
|
||||||
|
total := 0
|
||||||
|
for total < want {
|
||||||
|
n, err := conn.Read(buf[total:])
|
||||||
|
if err != nil {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
total += n
|
||||||
|
}
|
||||||
|
return buf[:total]
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleConnection_WritesHandshakeFirst(t *testing.T) {
|
||||||
|
client, server := newConnPair(t)
|
||||||
|
t.Cleanup(func() { closeQuietly(t, client, server) })
|
||||||
|
|
||||||
|
h := mysql.NewHandler(
|
||||||
|
&fakeTokenLookup{tokens: map[string]*token.Token{}},
|
||||||
|
&fakeEventRecorder{},
|
||||||
|
)
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
h.HandleConnection(context.Background(), server)
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
got := readAllAvailable(t, client, 80)
|
||||||
|
require.NotEmpty(t, got)
|
||||||
|
require.Equal(t, byte(0x00), got[3], "handshake sequence ID 0")
|
||||||
|
require.Equal(t, byte(0x0a), got[4], "protocol version 10")
|
||||||
|
require.Contains(t, string(got), "5.7.40-canary")
|
||||||
|
|
||||||
|
require.NoError(t, client.Close())
|
||||||
|
<-done
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleConnection_NonCanaryUsernameDropsSilently(t *testing.T) {
|
||||||
|
client, server := newConnPair(t)
|
||||||
|
t.Cleanup(func() { closeQuietly(t, client, server) })
|
||||||
|
|
||||||
|
tok := &token.Token{ID: "abc", Type: token.TypeMySQL}
|
||||||
|
lookup := &fakeTokenLookup{tokens: map[string]*token.Token{"abc": tok}}
|
||||||
|
rec := &fakeEventRecorder{}
|
||||||
|
h := mysql.NewHandler(lookup, rec)
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
h.HandleConnection(context.Background(), server)
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
_ = readAllAvailable(t, client, 80)
|
||||||
|
pkt := buildHandshakeResponse41ForHandler(t, "regular_user")
|
||||||
|
_, err := client.Write(pkt)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
<-done
|
||||||
|
require.Empty(
|
||||||
|
t,
|
||||||
|
rec.snapshot(),
|
||||||
|
"non-canary_-prefixed username must not produce an event",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleConnection_KnownTokenRecordsEventAndSendsErr(t *testing.T) {
|
||||||
|
client, server := newConnPair(t)
|
||||||
|
t.Cleanup(func() { closeQuietly(t, client, server) })
|
||||||
|
|
||||||
|
tok := &token.Token{ID: "mytoken", Type: token.TypeMySQL}
|
||||||
|
lookup := &fakeTokenLookup{
|
||||||
|
tokens: map[string]*token.Token{"mytoken": tok},
|
||||||
|
}
|
||||||
|
rec := &fakeEventRecorder{}
|
||||||
|
h := mysql.NewHandler(lookup, rec)
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
h.HandleConnection(context.Background(), server)
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
_ = readAllAvailable(t, client, 80)
|
||||||
|
pkt := buildHandshakeResponse41ForHandler(t, "canary_mytoken")
|
||||||
|
_, err := client.Write(pkt)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
errPkt := readAllAvailable(t, client, 200)
|
||||||
|
require.NotEmpty(t, errPkt)
|
||||||
|
require.Contains(
|
||||||
|
t,
|
||||||
|
string(errPkt),
|
||||||
|
`Access denied for user 'canary_mytoken'`,
|
||||||
|
)
|
||||||
|
require.Contains(t, string(errPkt), "28000")
|
||||||
|
|
||||||
|
<-done
|
||||||
|
|
||||||
|
events := rec.snapshot()
|
||||||
|
require.Len(t, events, 1)
|
||||||
|
require.Equal(t, "mytoken", events[0].TokenID)
|
||||||
|
require.NotEmpty(
|
||||||
|
t,
|
||||||
|
events[0].SourceIP,
|
||||||
|
"source IP must be captured from RemoteAddr",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleConnection_EventExtraContainsKubectlEquivalentFields(
|
||||||
|
t *testing.T,
|
||||||
|
) {
|
||||||
|
client, server := newConnPair(t)
|
||||||
|
t.Cleanup(func() { closeQuietly(t, client, server) })
|
||||||
|
|
||||||
|
tok := &token.Token{ID: "extra-probe", Type: token.TypeMySQL}
|
||||||
|
lookup := &fakeTokenLookup{
|
||||||
|
tokens: map[string]*token.Token{"extra-probe": tok},
|
||||||
|
}
|
||||||
|
rec := &fakeEventRecorder{}
|
||||||
|
h := mysql.NewHandler(lookup, rec)
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
h.HandleConnection(context.Background(), server)
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
_ = readAllAvailable(t, client, 80)
|
||||||
|
_, err := client.Write(
|
||||||
|
buildHandshakeResponse41ForHandler(t, "canary_extra-probe"),
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
_ = readAllAvailable(t, client, 200)
|
||||||
|
<-done
|
||||||
|
|
||||||
|
events := rec.snapshot()
|
||||||
|
require.Len(t, events, 1)
|
||||||
|
require.NotEmpty(t, events[0].Extra)
|
||||||
|
|
||||||
|
var extra map[string]any
|
||||||
|
require.NoError(t, json.Unmarshal(events[0].Extra, &extra))
|
||||||
|
require.Equal(t, "canary_extra-probe", extra["mysql_username"])
|
||||||
|
require.Equal(
|
||||||
|
t,
|
||||||
|
"0x0001a285",
|
||||||
|
extra["mysql_client_capabilities"],
|
||||||
|
"capabilities formatted as 0x + 8 hex digits",
|
||||||
|
)
|
||||||
|
charset, ok := extra["mysql_client_charset"].(float64)
|
||||||
|
require.True(t, ok, "charset must decode as a JSON number")
|
||||||
|
require.Equal(t, 0x21, int(charset))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleConnection_UnknownTokenDoesNotRecordOrErrAdvertise(
|
||||||
|
t *testing.T,
|
||||||
|
) {
|
||||||
|
client, server := newConnPair(t)
|
||||||
|
t.Cleanup(func() { closeQuietly(t, client, server) })
|
||||||
|
|
||||||
|
lookup := &fakeTokenLookup{tokens: map[string]*token.Token{}}
|
||||||
|
rec := &fakeEventRecorder{}
|
||||||
|
h := mysql.NewHandler(lookup, rec)
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
h.HandleConnection(context.Background(), server)
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
_ = readAllAvailable(t, client, 80)
|
||||||
|
_, err := client.Write(
|
||||||
|
buildHandshakeResponse41ForHandler(t, "canary_does-not-exist"),
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
<-done
|
||||||
|
|
||||||
|
require.Empty(t, rec.snapshot(), "unknown token must not record an event")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleConnection_LookupErrorIsSilent(t *testing.T) {
|
||||||
|
client, server := newConnPair(t)
|
||||||
|
t.Cleanup(func() { closeQuietly(t, client, server) })
|
||||||
|
|
||||||
|
lookup := &fakeTokenLookup{err: errors.New("db down")}
|
||||||
|
rec := &fakeEventRecorder{}
|
||||||
|
h := mysql.NewHandler(lookup, rec)
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
h.HandleConnection(context.Background(), server)
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
_ = readAllAvailable(t, client, 80)
|
||||||
|
_, err := client.Write(
|
||||||
|
buildHandshakeResponse41ForHandler(t, "canary_anything"),
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
<-done
|
||||||
|
|
||||||
|
require.Empty(t, rec.snapshot())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleConnection_NilEventRecorderTolerated(t *testing.T) {
|
||||||
|
client, server := newConnPair(t)
|
||||||
|
t.Cleanup(func() { closeQuietly(t, client, server) })
|
||||||
|
|
||||||
|
tok := &token.Token{ID: "noevents", Type: token.TypeMySQL}
|
||||||
|
lookup := &fakeTokenLookup{tokens: map[string]*token.Token{"noevents": tok}}
|
||||||
|
h := mysql.NewHandler(lookup, nil)
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
h.HandleConnection(context.Background(), server)
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
_ = readAllAvailable(t, client, 80)
|
||||||
|
_, err := client.Write(
|
||||||
|
buildHandshakeResponse41ForHandler(t, "canary_noevents"),
|
||||||
|
)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
errPkt := readAllAvailable(t, client, 200)
|
||||||
|
require.NotEmpty(
|
||||||
|
t,
|
||||||
|
errPkt,
|
||||||
|
"ERR packet still sent even without recorder wired",
|
||||||
|
)
|
||||||
|
<-done
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHandleConnection_BadAuthPacketDropsSilently(t *testing.T) {
|
||||||
|
client, server := newConnPair(t)
|
||||||
|
t.Cleanup(func() { closeQuietly(t, client, server) })
|
||||||
|
|
||||||
|
rec := &fakeEventRecorder{}
|
||||||
|
h := mysql.NewHandler(&fakeTokenLookup{}, rec)
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
h.HandleConnection(context.Background(), server)
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
_ = readAllAvailable(t, client, 80)
|
||||||
|
_, err := client.Write([]byte{0x00, 0x00, 0x00, 0x01})
|
||||||
|
require.NoError(t, err)
|
||||||
|
<-done
|
||||||
|
|
||||||
|
require.Empty(t, rec.snapshot())
|
||||||
|
}
|
||||||
|
|
@ -0,0 +1,84 @@
|
||||||
|
// ©AngelaMos | 2026
|
||||||
|
// server.go
|
||||||
|
|
||||||
|
package mysql
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
|
"net"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
acceptShutdownGrace = 2 * time.Second
|
||||||
|
tcpNetwork = "tcp"
|
||||||
|
)
|
||||||
|
|
||||||
|
type ConnectionHandler interface {
|
||||||
|
HandleConnection(ctx context.Context, conn net.Conn)
|
||||||
|
}
|
||||||
|
|
||||||
|
func Run(
|
||||||
|
ctx context.Context,
|
||||||
|
addr string,
|
||||||
|
h ConnectionHandler,
|
||||||
|
) error {
|
||||||
|
if h == nil {
|
||||||
|
return fmt.Errorf("mysql: nil connection handler")
|
||||||
|
}
|
||||||
|
var lc net.ListenConfig
|
||||||
|
listener, err := lc.Listen(ctx, tcpNetwork, addr)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("mysql: listen %s: %w", addr, err)
|
||||||
|
}
|
||||||
|
defer func() {
|
||||||
|
if cErr := listener.Close(); cErr != nil &&
|
||||||
|
!errors.Is(cErr, net.ErrClosed) {
|
||||||
|
slog.WarnContext(ctx, "mysql: close listener", "error", cErr)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
slog.InfoContext(
|
||||||
|
ctx,
|
||||||
|
"mysql: listener started",
|
||||||
|
"addr",
|
||||||
|
listener.Addr().String(),
|
||||||
|
)
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
<-ctx.Done()
|
||||||
|
if cErr := listener.Close(); cErr != nil &&
|
||||||
|
!errors.Is(cErr, net.ErrClosed) {
|
||||||
|
slog.WarnContext(
|
||||||
|
ctx,
|
||||||
|
"mysql: shutdown close listener",
|
||||||
|
"error",
|
||||||
|
cErr,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for {
|
||||||
|
conn, aErr := listener.Accept()
|
||||||
|
if aErr != nil {
|
||||||
|
if errors.Is(ctx.Err(), context.Canceled) ||
|
||||||
|
errors.Is(aErr, net.ErrClosed) {
|
||||||
|
wg.Wait()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
slog.WarnContext(ctx, "mysql: accept", "error", aErr)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
wg.Add(1)
|
||||||
|
go func(c net.Conn) {
|
||||||
|
defer wg.Done()
|
||||||
|
h.HandleConnection(ctx, c)
|
||||||
|
}(conn)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -0,0 +1,157 @@
|
||||||
|
// ©AngelaMos | 2026
|
||||||
|
// server_test.go
|
||||||
|
|
||||||
|
package mysql_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"net"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/token/generators/mysql"
|
||||||
|
)
|
||||||
|
|
||||||
|
type stubHandler struct {
|
||||||
|
called atomic.Int32
|
||||||
|
lastCloseErr atomic.Value
|
||||||
|
closed chan struct{}
|
||||||
|
once sync.Once
|
||||||
|
}
|
||||||
|
|
||||||
|
func newStubHandler() *stubHandler {
|
||||||
|
return &stubHandler{closed: make(chan struct{})}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *stubHandler) HandleConnection(_ context.Context, conn net.Conn) {
|
||||||
|
s.called.Add(1)
|
||||||
|
if err := conn.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
|
||||||
|
s.lastCloseErr.Store(err)
|
||||||
|
}
|
||||||
|
s.once.Do(func() { close(s.closed) })
|
||||||
|
}
|
||||||
|
|
||||||
|
func dialOnce(addr string) error {
|
||||||
|
conn, err := net.DialTimeout("tcp", addr, 2*time.Second)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if cErr := conn.Close(); cErr != nil && !errors.Is(cErr, net.ErrClosed) {
|
||||||
|
return cErr
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func runServer(
|
||||||
|
t *testing.T,
|
||||||
|
ctx context.Context,
|
||||||
|
h mysql.ConnectionHandler,
|
||||||
|
) (string, <-chan error) {
|
||||||
|
t.Helper()
|
||||||
|
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||||
|
require.NoError(t, err)
|
||||||
|
addr := listener.Addr().String()
|
||||||
|
require.NoError(t, listener.Close())
|
||||||
|
|
||||||
|
done := make(chan error, 1)
|
||||||
|
go func() {
|
||||||
|
done <- mysql.Run(ctx, addr, h)
|
||||||
|
}()
|
||||||
|
|
||||||
|
require.Eventually(t, func() bool {
|
||||||
|
return dialOnce(addr) == nil
|
||||||
|
}, 2*time.Second, 25*time.Millisecond, "server must accept on %s", addr)
|
||||||
|
|
||||||
|
return addr, done
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRun_NilHandlerReturnsError(t *testing.T) {
|
||||||
|
err := mysql.Run(context.Background(), "127.0.0.1:0", nil)
|
||||||
|
require.Error(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRun_AcceptsConnectionAndDispatchesToHandler(t *testing.T) {
|
||||||
|
h := newStubHandler()
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
addr, done := runServer(t, ctx, h)
|
||||||
|
before := h.called.Load()
|
||||||
|
|
||||||
|
require.NoError(t, dialOnce(addr))
|
||||||
|
|
||||||
|
require.Eventually(t, func() bool {
|
||||||
|
return h.called.Load() > before
|
||||||
|
}, 2*time.Second, 25*time.Millisecond,
|
||||||
|
"single dial must dispatch to handler")
|
||||||
|
|
||||||
|
cancel()
|
||||||
|
require.NoError(t, <-done)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRun_HandlesMultipleConcurrentConnections(t *testing.T) {
|
||||||
|
h := newStubHandler()
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
addr, done := runServer(t, ctx, h)
|
||||||
|
before := h.called.Load()
|
||||||
|
|
||||||
|
const n = 10
|
||||||
|
results := make(chan error, n)
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for range n {
|
||||||
|
wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
results <- dialOnce(addr)
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
close(results)
|
||||||
|
|
||||||
|
for err := range results {
|
||||||
|
require.NoError(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
require.Eventually(t, func() bool {
|
||||||
|
return h.called.Load() >= before+int32(n)
|
||||||
|
}, 2*time.Second, 25*time.Millisecond,
|
||||||
|
"all %d connections must be dispatched", n)
|
||||||
|
|
||||||
|
cancel()
|
||||||
|
require.NoError(t, <-done)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRun_ContextCancellationStopsListenerCleanly(t *testing.T) {
|
||||||
|
h := newStubHandler()
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
|
||||||
|
addr, done := runServer(t, ctx, h)
|
||||||
|
require.NoError(t, dialOnce(addr))
|
||||||
|
cancel()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case err := <-done:
|
||||||
|
require.NoError(t, err)
|
||||||
|
case <-time.After(2 * time.Second):
|
||||||
|
t.Fatal("Run did not return within 2s after context cancellation")
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := net.DialTimeout("tcp", addr, 100*time.Millisecond)
|
||||||
|
require.Error(t, err, "listener must be closed after context cancellation")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRun_InvalidAddressReturnsError(t *testing.T) {
|
||||||
|
err := mysql.Run(
|
||||||
|
context.Background(),
|
||||||
|
"not-a-valid-addr:not-a-port",
|
||||||
|
newStubHandler(),
|
||||||
|
)
|
||||||
|
require.Error(t, err)
|
||||||
|
}
|
||||||
Loading…
Reference in New Issue