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