From 924ff2a9d012aaa6164c8997421e1961d008e923 Mon Sep 17 00:00:00 2001 From: Lachlan Harris Date: Sat, 5 Sep 2026 20:55:11 +1000 Subject: [PATCH 1/4] chore: basic PCAP read stubs --- .github/workflows/ci.yaml | 10 +++++- cmd/pcap.go | 21 +++++++++++ go.mod | 1 + go.sum | 14 ++++++++ internal/capture/pcapng.go | 71 +++++++++++++++++++++++++++++++++++++ smtp.pcap | Bin 0 -> 27850 bytes 6 files changed, 116 insertions(+), 1 deletion(-) create mode 100644 cmd/pcap.go create mode 100644 internal/capture/pcapng.go create mode 100644 smtp.pcap diff --git a/.github/workflows/ci.yaml b/.github/workflows/ci.yaml index 22e52b7..90a3e27 100644 --- a/.github/workflows/ci.yaml +++ b/.github/workflows/ci.yaml @@ -22,6 +22,10 @@ jobs: runs-on: ${{ matrix.os }} steps: - uses: actions/checkout@v7 + - name: Install libpcap + run: | + sudo apt-get update + sudo apt-get install -y libpcap-dev - name: Install Go uses: actions/setup-go@v7 with: @@ -33,7 +37,7 @@ jobs: test: strategy: matrix: - os: [windows-latest, ubuntu-latest, ubuntu-24.04-arm] + os: [ubuntu-latest] # [windows-latest, ubuntu-latest, ubuntu-24.04-arm] -- skipped due to libpcap requirement name: test runs-on: ${{ matrix.os }} steps: @@ -48,6 +52,10 @@ jobs: shell: bash if: ${{ steps.setup-go.outputs.cache-hit != 'true' }} run: go mod download + - name: Install libpcap + run: | + sudo apt-get update + sudo apt-get install -y libpcap-dev - name: go build run: go build - name: go test diff --git a/cmd/pcap.go b/cmd/pcap.go new file mode 100644 index 0000000..ed592d2 --- /dev/null +++ b/cmd/pcap.go @@ -0,0 +1,21 @@ +package cmd + +import ( + "github.com/spf13/cobra" + + "github.com/lachlanharrisdev/gonetsim/internal/capture" +) + +var pcapCmd = &cobra.Command{ + Use: "pcap ", + Short: "Open a PCAP file", + Args: cobra.ExactArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + err := capture.ReadPcap(args[0]) + return err + }, +} + +func init() { + rootCmd.AddCommand(pcapCmd) +} diff --git a/go.mod b/go.mod index 2e002ae..3228721 100644 --- a/go.mod +++ b/go.mod @@ -19,6 +19,7 @@ require ( require ( github.com/fsnotify/fsnotify v1.10.1 // indirect github.com/go-viper/mapstructure/v2 v2.5.0 // indirect + github.com/google/gopacket v1.1.19 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/knadh/koanf/maps v0.1.3 // indirect github.com/mitchellh/copystructure v1.2.0 // indirect diff --git a/go.sum b/go.sum index 73cfc15..d469aad 100644 --- a/go.sum +++ b/go.sum @@ -7,6 +7,8 @@ github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx5 github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo= github.com/go-viper/mapstructure/v2 v2.5.0 h1:vM5IJoUAy3d7zRSVtIwQgBj7BiWtMPfmPEgAXnvj1Ro= github.com/go-viper/mapstructure/v2 v2.5.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM= +github.com/google/gopacket v1.1.19 h1:ves8RnFZPGiFnTS0uPQStjwru6uO6h+nlr9j6fL7kF8= +github.com/google/gopacket v1.1.19/go.mod h1:iJ8V8n6KS+z2U1A8pUwu8bW5SyEMkXJB8Yo/Vo+TKTo= github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= github.com/knadh/koanf/maps v0.1.3 h1:P1z7EvTqdFBrPYbzSvorvrpib+sjkUMxf0FVvA5NKK4= @@ -46,12 +48,24 @@ github.com/stretchr/testify v1.9.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8 github.com/yuin/gopher-lua v1.1.2 h1:yF/FjE3hD65tBbt0VXLE13HWS9h34fdzJmrWRXwobGA= github.com/yuin/gopher-lua v1.1.2/go.mod h1:7aRmXIWl37SqRf0koeyylBEzJ+aPt8A+mmkQ4f1ntR8= go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= +golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= +golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI= +golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY= +golang.org/x/mod v0.1.1-0.20191105210325-c90efee705ee/go.mod h1:QqPTAvyqsEbceGzBzNggFXnrqF1CaUcvgkdR5Ot7KZg= +golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= +golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To= golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU= +golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM= golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= +golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= +golang.org/x/tools v0.0.0-20200130002326-2f3ba24bd6e7/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28= +golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/internal/capture/pcapng.go b/internal/capture/pcapng.go new file mode 100644 index 0000000..a197a63 --- /dev/null +++ b/internal/capture/pcapng.go @@ -0,0 +1,71 @@ +package capture + +import ( + "encoding/binary" + "fmt" + "log" + "os" + + "github.com/google/gopacket" + "github.com/google/gopacket/pcap" + "github.com/google/gopacket/pcapgo" +) + +func ReadPcap(pcapFile string) error { + f, err := os.Open(pcapFile) + if err != nil { + return err + } + defer f.Close() + + buf := make([]byte, 4) + _, err = f.ReadAt(buf, 0) + if err != nil { + return err + } + + if binary.BigEndian.Uint32(buf) == 0x0a0d0d0a { + // pcapng + return ReadPcapNG(f) + } + magic := binary.BigEndian.Uint32(buf) + littleMagic := binary.LittleEndian.Uint32(buf) + + if magic == 0xa1b2c3d4 || magic == 0xa1b23c4d || + littleMagic == 0xa1b2c3d4 || littleMagic == 0xa1b23c4d { + //pcap + return ReadLegacyPcap(pcapFile) + } + + return fmt.Errorf("unknown pcap header") +} + +func ReadPcapNG(f *os.File) error { + reader, err := pcapgo.NewNgReader(f, pcapgo.DefaultNgReaderOptions) + if err != nil { + return err + } + + packetSource := gopacket.NewPacketSource(reader, reader.LinkType()) + for packet := range packetSource.Packets() { + fmt.Println(packet) + } + + return nil +} + +func ReadLegacyPcap(f string) error { + handle, err := pcap.OpenOffline(f) + if err != nil { + log.Fatal(err) + } + defer handle.Close() + + // Loop through packets in file + packetSource := gopacket.NewPacketSource(handle, handle.LinkType()) + for packet := range packetSource.Packets() { + fmt.Println(packet) + } + + return nil +} diff --git a/smtp.pcap b/smtp.pcap new file mode 100644 index 0000000000000000000000000000000000000000..931b43b3b878dc559b7423df26363bfe75c6a512 GIT binary patch literal 27850 zcmeHw3y@padEO-{>aq6Pq^EIdM;)Do#Arz^u)w}Zf{zdjU?0>zpj|8}$Fsy#Grk7^ ze09Qq&xUc98NAQ>?m^{uN*%g#?yCn4BRRLQtY%Xh z+{2%^-Z%wFeDzRrc<9@tWa8;gN@grj^0TP}?|l9X5Aw0!`SPLXiQ1D-TyOk)eCDeM z*3&~T3aneVg@sj#yc4T3x$~;bS;e9Qyw!UmV`HJn6kZeKp~&R1dqI3Ww>p1!QcaFc zj2%B74(qQ{^Qq**{L=i=to|HZOV6p)LTr9XEiBK@FL@89V=L+O!m4~ybIApdG#~yK z*BjHA312;wo;ZAqQ2GzLDE-+_zmj=1b?Eiq|K);V$avmJB+$#>88qeqjISO7MTcGz z0)4EDK(B(>e@M2?LX(Y;PpI>EE)}z-$wKDB+LT_zTVEbD(g4X<52iDRj|is!ybIGm z{PZhX1DJj$@EeA4_xSm(?Ci`|Zg!?|$J}%wH+Qa*8DF`)P%7332fhg03*Xq+5BKi~ z?jP;MJ!=5>7m54XNsas3?9Aotcy1$Gx&ycq?kB(52De*HM-PVt_uuKl{YBvZ4&eT= z^wWlsojYGQOEdNziR9K|?^t(Bm!N)M0 z$*@|!T|=9G>}|%qK$))|yq6Cj7rg&rC*J9I1MlO&yYb}}!&r%@(ki_StgGhz7ju?V zgjO80%V)RDg5yAoYSjM(3;gwOUSmMv%BzRzr-oh@)Vr8IdmO08i24g}*Qm!p_zLvn z-~eGh_rf(Tt<}vmp~I(T0#9{KAchIdU;_4;7lBPI9n-Krb7;VTWtUeEe)kQ1SMaNL zk#+|7%@V)U889}?K(bu3s;XqUuDNcFD65>a%j>FP*9s~S3MlZ7sfu>lQkgAf*IcWZ z*P~zeYXio2_Z#x+UW)x+KfKJDeE%=HW-^PJJdir@dgLn~FpQb1Q;Mpgf)3AAHF4Ps!GRkxPQ9Ent6Qf7LHs#=amm`OrEteZjL6r-#c(7Zo0yyq zhr%Z&r%sH|gad;EcLo<>L9HrgVP~t3>*Q-{IiI()mRhdYijH$hMaDROx=^*uTyVae zvo@nDmaWXJTcs1z8~5LTcQRqEL~bt}w{I(~ZJl?*F=`qg-vzV0UdPHuvGT#XJ2u#%`8nv)J*Eqikd0k zc0tX}#8YaneAhX3Zuyw2Zm*Z7)a_ff3u<8|JFgb*U#p#roXoCI7hS$FWhbza|y#V2BSc|YYk=5cb>aSLYT;zs1;v6JtM z$718hW+qPrBgbJ@Cnn>^g2yK(!ogTL)V{?nKh4tx5Y<41cve%UY7vWz-%`T=-t~I8YVl#Kd)?zV` zSE#7zXXm2ojiEL;z;%SHg@0y?hkW5wc3F}0RNkqU%o^OP!2$Bp8|0?ens@W4S%eBG z!%%EkSK-tMIdyVdPTc{Uou25%rUr2strShWOcb&Ov+7#4(}B7hG~KLiLlk_DP1F0# zTh(B)oOP(IqUu=2u5o!-fve{{$&@-@NK8@%CEb8SPHICP;_t3*dfy%~m@Gh~=Es-GK$;hq?cjEy>__8GT5 z-DosIjq#9ET^~)ajIJc(K~ixd9Dr{^l0sqP-KoW7I;Q5*=~OUz$J+dP0%-CYw0S_q zmzUDXr8F-Ho65S9#6FvjCz^apfF-MDs?A9E3e**Gor;<0-ecRkZkKC>R7&j~-Ps0HP9B-Y(np>s)no6c^6s@pUk7X2B~r zywCekvWr{MK-?^4sX%~~NqW~`WH$F+iy=$M*NlWH8V2{nbE_>_B5MqrHFzRR0F6$JIbNWg{i2)b6)&IdSKGYa44(ui_F77jg! z1+Vfe8Qpb)vILYx1jV(j=!p{|&`!~*Ml;2_=u7fMFlS{QtTP<~*m;3r=>axuhUtip zd&Zb?5EwS-8e*pX$Ev_dDqejw86wG`TZkVy)(N^s=L8dUow8Y^2v)!fO#=2O79eG1 z+6|~qr@={cej!pVUZ{X4O*lQorS(+Umr;XHG#!;vRL%iNcMj zAU={wxN0Mc=}}T4jI{t(&d!z0O@0vxk4**Cs6gh^r-t4`mpOkgMx&TsWrU)+<!AX0%4R13}Zc7#6dL4peAqyMjHWVX9m?)!yu#;u%u19$J_2E5VHws95S=%YUyqVPMm13ZfL(6kuk zT6%A3gXHbc;H1!qZ0ptR6o9ljT@u!GI%#6z0!zsYtD)Lvt-puX>(lrqHen-`2n{Kr zRLdI^tpsnaXOO`GFYrQ&^924vg&c)-D{mp?xDE$z(?Sq)RLz$vRc8bCloISzIq}Zt z5HrC<)hm+6ghve-sab9faaMV~0b^aNFz)iRo|(E`%n9J>)+n7IkB3eWmU0VDy_iF)x1p+*%-t@#HM5A>V+9n5 z29LAksyu>ol`^xJKw5uVoQGRbZnkP!NDBaGCql)kIGv>TAX}h*?-F$r52k9n zNzKON>Wmr>jq!5A;u>$5aA<%ff~2bD%>1%_y}h zn8&G*sBZN17`RQTAY7s64E=>h- z7ycio=D*Ni^RI~J|9+?Dx3a$G|JK8XF@jL5P~>ie{u!)g>?*t>@6$x{-DX8*x-$7O}G%k%90UcBkE(#@-~^y%;C`0LIS^H?B$o9$lPxk38Q!0)RxzrV zR&y0{K`wdI_S6z@>7z*|d0k65r5Xy-wKH}R9w_%;N(g)Nc!w(+CapNOXw0Tq_+o&W zh!DVmURT3ad%XZ2$7anXozrVj9$F7knvvY2AOsdQe1RR1*(bDQC9I8L{I=WFj9rDM zOxq=>tL1XuMgmyt0PP!&ko;9wWzDjPXBL@NY$O4pj8$t;chrR#AP{9xd<~VUIt@If z*OU5_D!6EgR?1sutzNZUq&%*!NGAHyPV>PQhQSJEi**c?vnl^hbxTB!R2BUVblXtK zLj0)HRR))lYJ*1))`dpW6v+!kph&B=dW8ry?T^6f2E;pBI=NvNi$W+ePaPJ+X_SkO zN&OD^Df1FCLS=ubd#IG*>bj(76Lm-r7SPA0rC|7A&f3UUDj|FsnvTV9UrTZHX?vYA zfgvzpn>3J95xK&MxNsER8kwh2VA_O+6rbeboo(PauP-;Tg+V5()W|BZ4{fZHv(2Kj z4!kwByD3DT#sp!QgwH`)W>0gTit?9^`J5J?VQJQNxNYTxRho;?qBUfv&Xmg4`A94d zt+_TO9uEqSBo-FjD3+h8Bl3i)s`WB#`8w?<6%OI^QG8@R+A{=EwFWWB)?HY0xGa|t z#uu%SOiAC~L=&1r@RV7^7$u5GajZ55_Vae#q?Ife)-+k(u&Yj)?gVM+X(^9{!&qAo zmof^v!SkqAZTJ>U^K0o01hdrqw?8&uEbTSq)wRp;%-*xfcVRy#s zy!`!NN*a0(;DKv@u;U)UP2btw19&s<0g$^s0PcDLyGVTX>>_>X@T&L~KkD==(#>6@ z)#xq^cNb}!4K9YdySqq|rrq5|`hRN|i5h?QV}1RJcZ)wU(()(vf3E3IU}p*b z#Ev^lU-_loouxPP&JwxncR${jyMA7{>$Vnm-E1|v3)yRM*N)liEfj3+X0LZn%K!hK zy(V`>AMeXuHwt$hZE@EZ&NaCU88UF!jv2C}ukL2Z-pm=YKLvMv|HFG)wSnD*pTk4n zlX|{Sb=C7d@W8Woy_Pxoq@L`{L_{ z(N$HlTdSdVYc*QM86(J=?$&B-Ul+1ltFc?FAr&J3$7(gG@qeMe#$OYS|HDp=fAQnK z#{cQdhS6Q|vFq)=dA;4v`qSN7jon%eX?BpBo8zI~T8;nPwHnm;UpsnTzj)yli5Gse zH(t2-@zjCfeV_Uvs zeeBkK?ACoqE#7Y32P+qvRX)1JX1DGG^(p_G>prOQr*7-3@h^zR@2cIn|0ll2|EnJv z#?`7v0$kh@>yl6oi6S=nZ@X(krMCA1%PytliUl^NLh;Rr z)CBouSf~wRX)HQCP}0is7yxE{D=WEB_TtsxphE*)GpKh${e1-`_|4L(-eILQI9@=_ zJZ|@klZh*tl~uuheP?q@QlORcihaB~8<}OZ0;ymXb^BT0`l(~OhIqwtoq9Eknk!Zj zw%56`Vi!oExQDZA-iK~YOR9pq3!ytHdMtEYS3aTYtx<4L)TMiA{8C_4Xw|AtQRPv6 zCY4_-NJWfQn?x|A7+>0343P?FxK%JNtUjr`p;#eNwguju;AuG&k#@(e3r zs21pVvSc)QaU*oZS~bDym{< zrTP)9xywgTkAqd*AdYQT|&xxVB?RWNb!Vj?-cHhb|RQBAUh&M%YG z5)BfmQ{Gn>S^t0LBKog3K<7|FcW{zV(J4yYFbYc7bdk4$TjUHA@L{#e3d*3apwwlo zEcDd{@0xuA$xUglEpn3UD6Mn@l-@L&1KiQ}0zJQmc$jtHb+q?MfoEVD$jAcdLR=4M z>stcz$)}$paaaM5vOk=F;ju9_{)LD8#tYwR5#{uAs_8OGna^Wawr z%L1%E@xq;AJ-T-4D^G~^IP`iXc)&1JK)BojJ=iIcF{MjjhVAm0kR{z|w%3)=yCmQ& z^fp`qUrF(PQQ9jsaaf2iD@p>OtG!m;LJK>zj4z=oCJX@~2*U74L1)gXvua$rR7971 zHM7|~jT5f0xt^tw8@i2xbX^7u@tm7x(@D3kl~Pz?J2NMTD)fjDE3vEAd+X>s8NnOr z{|E#@5v}`l69za>w-a#CcvC2aouEco5s!jlk)`dm=V}4Qum8V8=M13h2Fd_>6skcoDiIFMdnMR@#VRLCsRlAEqwTz#s?Mvz zjrt0uSB439)uLslp=7nXuY3G<3{BQ}@AnWTzbcpVKI~dBkXT=@QRe2U$%$hV6QZ#@ zNAQgTd&1+FXEw-5vqp#(ZPX~PR208m6 zZAB1msF(y>CDt(l%H*NgE7EiU60>dV7>}$#2{w<(KsT!6!YMiqW27 z78RatEdxjHBUeH?T@BduK+7M_cp$59eqfjpO< z&dM;+aUnm#3=||Zx(SXRK=^p*0@8jFOs!R>|AFIjGy;Jm7)!br(CbTh1sWf`uQtU| zd$95It5_j%xkq?i-Fhoxs!QaaMuGi3;O_$L4)`4tM}sn$J`@u9({_OUiM~J0wfE5D znzi2`e9?9dK?!>NFlk5(8Nip|7O*^fKxkkI3a24R*h3;fpLQzc9u=;E>2@HD*{jA= zo*a2@8KI;IE$&{ROT$9~ZQ=u>fpls*a`$+~6{s(wZe?wm2R!XBFxx#ZqJ4e*w#MKO- zrNxF~qMqe@gh|aPSR~J%JcdLMkM4GAxdqzpqQ#H>Gon1r%#mgr71}g_7yNBkyj)Rt zV2x(5rk~bLQW~WIey9Z;+5pAY-niyt;mJ~8eKVy)(srJ4Ewc)*9auJ_THUTl3bEys zol$VE*>^YWyQ4rIO`2@r&!&WCGAAD<$J)*}d_Hck|8=+}}% zqXC-?$4x5-t_U6=lW_8$Ge^e*kWA?>#5GBcx3?cp=|H@_dlzX91qp|N)3Q@=r1@$@ z`={gvY2+hk>Gm$_d^REx1+7oK_AgNobe1|B>$FdaXj=S1yA1D@=Taa}C}m)r#M{$E z0b11fH~+i0_S;{0RpNzD_QnhM|0H!Fao~gBWxVkFuN?^xRsyj4#0z(e_2{aOdtdX+ zm>XWh37<5qjEzX=TcgoiJ;Xs}>KJ+iaQSc@%Sh890fg=!6dNYi=O}d=2zRvE{fXdi zy-r;SRg37mbt{A>d&#Qd+x0?`&69X6q{N*k23hzIVo$&sCPH<&nRzL&J{H7(8q1{> z2)fj802o3+c&56?xCz)#e<0P8Yq|=U1+=GlbM%9{lT=fB56;(dejj)XeQjIB`n!sH zpqVLQl+m`>zW1cL&JR35$AF*Ra8<(Zkd)Xp(C%VkuAUYVn6%1mlt*(RGu@*4cJyeb zbRs6HO}v-npnE2V+~Cle%b*L0fdz(GgWi6dqx2RA(R$aKcFx({LJy!mwuOdGFD>kTVAuS-}-^%jU_s1w>84&TOgjezD#I^;T;Llv^I;xO%GZGv)U8F^L- zGI&}Tr;d(P*qwqG_H|3%WJs2R2zbH*DtWx?NdnRh47v5m6QWk!5L-WS1{6#O;x%be zC&@rQY8^0`;`cl!0=y1@yy!hIxbd?jJn{bqxQqRoK$ONd#B&od2XPkxu%Ufbg3Ovo zk&5Bs`50a@c^GY@I0RxH1{}V&J`Mt|1sO$B9B@C6pwAwT5JV90*6j^a1sD({0BRNN zO`fdKoDlp~3DBwWFZS2??-7mvO6QL9Lr?ngLV3+Fs(KrwAEoJDbSRu7UXN|(x15IT zE`c`{8n25uc?y#U|Fe0R=PF=+PPf?zLD4`OgeRbNxCDBOiU$^r5E8!v*E6jLClVez zdU|OM02?PA(7yvzuo#X%s^c63#tgW91AVqQV#s2!3I9eGstaFQc>)(|<1j`#BN_o6 z(7~t8SRBH`ocCgm6k6dWeAL42v6A3@gl-vQ3M6ZR8_p&HTh8aThhc3(j}?(IXbuZ_ zBoNUHc&aN@2#mmYLzgworH>M{PAoM)#E?f13wzVp9Bpj&+zp5hP7hV8v8h#+JEZPH z*aL8ylseJvg*f^=6!_80+<5?_dJYX}Cn~2sd)v&IVZNPS561LEBw74<99qrnG!_O- z4^gL9$T&DAL;O9`Q+7HS9CKp{qYcyf!(krzJmi9F2$-M}4U|yF1oAk_0{X%WTRhQ# zAi>*$Nf0Y8Xfa`cqo=?hIXhRrl@W)YIzX~)6p1*Qy<+I1S6@_yY+Z2}NkwppI%Wxr zls$F>bihOcx;T#<$dcQ1^op~^We))u=_kE}2yp65noerNw1)h9 zRV}4Qcw~84%=4WoeX2xYXCQmoBEPhfDE+3pF9lT6({1NyiJY& zxkH2f;)NebyztL@k+>N zyVI0h&ot$5$PPV!QWR7Ym5rEzQ|(?p@jhC!Z<%;u+@R+M`BsMeaa@^Ua|37YiCGtO zEvK#MY$XMvDX{+Ye1<&dMVGjxO;Ys=_)adfL0I1k6s`_*Hf+b+76Ik`tdi_30S&;! zxB8rB=r~Cno!o{?#ffSUzKLQ$eAuL%*GE@Y^cEE5TzYXK(2STwk!T79(m2fs;|WKb zNHOT&5g>O+q~srq0Y=pu+P+(fLjgfMtBO7Xhk6hh;|30k0hh|5QSfH~M5%dy{~2~z zM)9AAA7vJ)!g#hF+<2bsJw5p7IJvSQCbIzWNofK&v<}z!vwC%fe$Iy(4^vY0LlsQU_0t z)=dl==O)~AvzL3ydAf~J593;`3?QN+*m2Ks0sWExlEf)g0zP#MDm`8~>IEba;?=UbQT~S84GD#n@L=c;e%OS z(xRHjdu7$BHcQhoOp-*$4x~62SRaj(yM;~c%~n~ATCY0#?4VTJaZg-NB%H+2Yo)k2 zVeXpJm(yE4varIs)c?(>uT#BEgd%OLg-={B5`za1k;+JrDQ3$;V#PpSfESCdbE>?R zGODFv6uUh7BqDL5V6lh2?NM3awgR6fI@=sp$g>Bj@xK9$zX}V$SNBOEM)L5z)cCJ$ zb!z-auJ}93-wYbYc1Nh|BM(6w4l!HpYoPu@3{>~;B0z)P#&YVw-OpGr5VJpye^-)Y z&KLW_`b&ZJM6a=X-%nEq@`uvjAgq5p_wwAAfY)n-`}!Hq9}290+6(KUCsPOBTZ#V{ z!us~+{+9^r-TA(-UK3c4_rm(%YU;q>{iCt35Y`W$`CgE)PCeTf){g|%CwgH$awT=( zLu%+x2y5=;^-mJkw?E$()?W#%kN3iQ^hSa8uL*1LS0DX#!us-a`-~NQ=Bw*SqWg@* z(2oVy2eyTE=tJtEuNmI&lu&l=bYu~cuoHxvIi6S#u9XJ|IKo@-n+>J_oUa~c96o%X zFylYyWX4Bt1T%gFzue&Ge)|c-_zIvr{p4F(>)QLv9xQ?Kp>3fYdj042f!}Z*b$W1s zus-oyZ|R3oUSK`i1uGyJ9Rx-nHH<$4l%GDf$H-$!eDzT7&Y_&J{m}Y za;IiMxv_}a>DLdjkbv3(O0bMT0!C$UfTOL12Mpw=<#k2=7<+&574x9{O0Acs-f!U7b-j`UMt=*hE5^h4`^A$#e{k=E z-{5!n^lwoNlUz!kn@Pr(5{YCoF`Jl4%p|9i(}`H(ToQlcF+ErY0w2vG9pR!q~rO-+_I`+xM*fmhnFUQ%`SL literal 0 HcmV?d00001 From 3023e4e748ccadcc41f1919e8e3cd482a2fcac9d Mon Sep 17 00:00:00 2001 From: Lachlan Harris Date: Sun, 6 Sep 2026 17:26:42 +1000 Subject: [PATCH 2/4] feat: true pcap implementation --- .github/workflows/ci.yaml | 10 +- README.md | 15 + cmd/check.go | 92 ++++-- cmd/check_test.go | 28 ++ cmd/pcap.go | 97 +++++- cmd/pcap_test.go | 113 +++++++ cmd/run.go | 43 ++- cmd/script.go | 3 +- cmd/servicecfg.go | 32 +- cmd/targets.go | 127 ++++---- cmd/targets_test.go | 51 +-- docker/docker-compose.yml | 6 +- examples/handlers/ftp.lua | 2 +- examples/handlers/irc.lua | 2 +- examples/handlers/smtp.lua | 6 +- go.mod | 2 +- internal/capture/capture.go | 73 ----- internal/capture/capture_test.go | 82 ----- internal/capture/inspect.go | 100 ++++++ internal/capture/pcapng.go | 71 ---- internal/capture/recorder.go | 203 ++++++++++++ internal/capture/run.go | 188 +++++++++++ internal/capture/session.go | 303 +++++++++++++++++ internal/capture/session_test.go | 451 ++++++++++++++++++++++++++ internal/config/config.go | 86 +++-- internal/config/config_test.go | 37 ++- internal/config/default_config.toml | 11 +- internal/config/size.go | 41 +++ internal/config/size_test.go | 31 ++ internal/dnsserver/capture_test.go | 215 ++++++++++++ internal/dnsserver/config.go | 36 +- internal/dnsserver/dns_test.go | 244 +++++--------- internal/dnsserver/server.go | 56 +++- internal/handler/echo.go | 2 - internal/handler/handler.go | 10 +- internal/handler/handler_test.go | 298 +++++------------ internal/handler/lua.go | 348 -------------------- internal/handler/luabindings.go | 243 ++++++++++++++ internal/handler/luaconn.go | 123 +++++++ internal/handler/luapack_test.go | 170 ---------- internal/handler/sink.go | 6 +- internal/handler/sleep_sni_test.go | 71 ---- internal/handler/testdata/capture.lua | 6 - internal/handler/testdata/comment.lua | 5 + internal/httpserver/capture_test.go | 242 ++++++++++++++ internal/httpserver/config.go | 31 +- internal/httpserver/http_test.go | 151 +++------ internal/httpserver/server.go | 49 ++- internal/listener/config.go | 17 +- internal/listener/listener_test.go | 234 +++++++------ internal/listener/service.go | 84 +---- internal/listener/tcp.go | 77 +++-- internal/listener/udp.go | 60 ++-- internal/netx/netx.go | 107 ++++++ internal/observability/logging.go | 11 +- internal/service/manager_test.go | 13 +- internal/state/state.go | 37 --- internal/state/state_test.go | 36 +- internal/testutil/testutil.go | 35 ++ internal/tlsprovider/config.go | 7 +- internal/tlsprovider/tls_test.go | 8 + smtp.pcap | Bin 27850 -> 0 bytes 62 files changed, 3413 insertions(+), 1925 deletions(-) create mode 100644 cmd/check_test.go create mode 100644 cmd/pcap_test.go delete mode 100644 internal/capture/capture.go delete mode 100644 internal/capture/capture_test.go create mode 100644 internal/capture/inspect.go delete mode 100644 internal/capture/pcapng.go create mode 100644 internal/capture/recorder.go create mode 100644 internal/capture/run.go create mode 100644 internal/capture/session.go create mode 100644 internal/capture/session_test.go create mode 100644 internal/config/size.go create mode 100644 internal/config/size_test.go create mode 100644 internal/dnsserver/capture_test.go create mode 100644 internal/handler/luabindings.go create mode 100644 internal/handler/luaconn.go delete mode 100644 internal/handler/luapack_test.go delete mode 100644 internal/handler/sleep_sni_test.go delete mode 100644 internal/handler/testdata/capture.lua create mode 100644 internal/handler/testdata/comment.lua create mode 100644 internal/httpserver/capture_test.go create mode 100644 internal/netx/netx.go create mode 100644 internal/testutil/testutil.go delete mode 100644 smtp.pcap diff --git a/.github/workflows/ci.yaml b/.github/workflows/ci.yaml index 90a3e27..22e52b7 100644 --- a/.github/workflows/ci.yaml +++ b/.github/workflows/ci.yaml @@ -22,10 +22,6 @@ jobs: runs-on: ${{ matrix.os }} steps: - uses: actions/checkout@v7 - - name: Install libpcap - run: | - sudo apt-get update - sudo apt-get install -y libpcap-dev - name: Install Go uses: actions/setup-go@v7 with: @@ -37,7 +33,7 @@ jobs: test: strategy: matrix: - os: [ubuntu-latest] # [windows-latest, ubuntu-latest, ubuntu-24.04-arm] -- skipped due to libpcap requirement + os: [windows-latest, ubuntu-latest, ubuntu-24.04-arm] name: test runs-on: ${{ matrix.os }} steps: @@ -52,10 +48,6 @@ jobs: shell: bash if: ${{ steps.setup-go.outputs.cache-hit != 'true' }} run: go mod download - - name: Install libpcap - run: | - sudo apt-get update - sudo apt-get install -y libpcap-dev - name: go build run: go build - name: go test diff --git a/README.md b/README.md index 38c7b3b..5121474 100644 --- a/README.md +++ b/README.md @@ -98,6 +98,7 @@ name = "irc" type = "tcp" listen = ":6667" handler = "lua:handlers/irc.lua" +capture = true ``` Run it with `gonetsim run irc`, or skip using a pre-defined config entirely with `gonetsim run lua:handlers/irc.lua@:6667`. @@ -106,6 +107,20 @@ The [`examples/`](examples/) directory has a full sample config plus example IRC
+## Captures + +Every run saves everything it handles to a single pcapng file, typically `~/.local/share/gonetsim/runs/.pcapng`. GoNetSim prints the path on startup and a packet count on shutdown. Lua handlers can annotate interesting packets with `capture:comment("...")`, which shows up as a packet comment in Wireshark. + +```sh +gonetsim run http --output ./case.pcapng # choose the capture location +gonetsim pcap ./case.pcapng # summarize a capture +gonetsim check # also verifies captures can be written +``` + +Two things to know when reading captures: handshakes are synthesized (sequence numbers start at 0, Ethernet MACs are fake, timestamps mark when GoNetSim wrote the frame), and TLS services capture ciphertext, not plaintext. + +
+ ## Docker A lightweight distroless container setup lives in `docker/` and is built/published with `ko`. This is the recommended installation method if you require long periods of uptime, or if your system is incompatible with the provided binaries. diff --git a/cmd/check.go b/cmd/check.go index aae4233..0a140d9 100644 --- a/cmd/check.go +++ b/cmd/check.go @@ -4,19 +4,38 @@ import ( "errors" "fmt" "net" + "os" "path/filepath" - "strconv" "strings" "syscall" + "github.com/lachlanharrisdev/gonetsim/internal/capture" appconfig "github.com/lachlanharrisdev/gonetsim/internal/config" "github.com/lachlanharrisdev/gonetsim/internal/handler" + "github.com/lachlanharrisdev/gonetsim/internal/netx" "github.com/spf13/cobra" ) +func checkRunDir() error { + dir, err := capture.DefaultRunsDir() + if err != nil { + return fmt.Errorf("runs dir: %w", err) + } + if err := os.MkdirAll(dir, 0o755); err != nil { + return fmt.Errorf("runs dir %q: %w", dir, err) + } + f, err := os.CreateTemp(dir, ".writetest-*") + if err != nil { + return fmt.Errorf("runs dir %q is not writable: %w", dir, err) + } + _ = f.Close() + _ = os.Remove(f.Name()) + return nil +} + var checkCmd = &cobra.Command{ Use: "check", - Short: "Validate configuration and check enabled services can bind their ports", + Short: "Validate configuration, runs directory, and check enabled services can bind their ports", Args: cobra.NoArgs, RunE: func(cmd *cobra.Command, args []string) error { cfgRes, err := appconfig.LoadOrCreate(rootConfigPath) @@ -49,39 +68,44 @@ var checkCmd = &cobra.Command{ checks := []struct { name string enabled bool - run func() error + run func() (bool, error) binds []bindTarget }{ { name: "dns", enabled: cfg.DNS.Enabled, - run: func() error { - _, err := dnsConfig(cfg.DNS) - return err + run: func() (bool, error) { + conf, err := dnsConfig(cfg.DNS) + return conf.Capture, err }, binds: dnsBindTargets(cfg.DNS.Listen, cfg.DNS.Network), }, { name: "http", enabled: cfg.HTTP.Enabled, - run: func() error { - _, err := httpConfig(cfg.HTTP) - return err + run: func() (bool, error) { + conf, err := httpConfig(cfg.HTTP) + return conf.Capture, err }, binds: []bindTarget{{net: "tcp", addr: cfg.HTTP.Listen}}, }, { name: "https", enabled: cfg.HTTPS.Enabled, - run: func() error { - _, err := httpsConfig(cfg.HTTPS, configDir) - return err + run: func() (bool, error) { + conf, err := httpsConfig(cfg.HTTPS, configDir) + return conf.Capture, err }, binds: []bindTarget{{net: "tcp", addr: cfg.HTTPS.Listen}}, }, } var failures []string + fail := func(name string, err error) error { + failures = append(failures, err.Error()) + return write("%-8s FAIL %v\n", name, err) + } + captureWanted := false for _, c := range checks { if !c.enabled { if err := write("%-8s disabled\n", c.name); err != nil { @@ -89,16 +113,16 @@ var checkCmd = &cobra.Command{ } continue } - if err := c.run(); err != nil { - failures = append(failures, err.Error()) - if werr := write("%-8s FAIL %v\n", c.name, err); werr != nil { + capturing, err := c.run() + if err != nil { + if werr := fail(c.name, err); werr != nil { return werr } continue } + captureWanted = captureWanted || capturing if err := preflightBinds(c.binds); err != nil { - failures = append(failures, err.Error()) - if werr := write("%-8s FAIL %v\n", c.name, err); werr != nil { + if werr := fail(c.name, err); werr != nil { return werr } continue @@ -118,8 +142,7 @@ var checkCmd = &cobra.Command{ conf, err := listenerConfig(l, configDir) if err != nil { - failures = append(failures, err.Error()) - if werr := write("%-8s FAIL %v\n", l.Name, err); werr != nil { + if werr := fail(l.Name, err); werr != nil { return werr } continue @@ -127,26 +150,35 @@ var checkCmd = &cobra.Command{ // compile the lua script to catch errors if _, err := handler.New(conf.HandlerSpec, conf.BaseDir, nil); err != nil { - failures = append(failures, err.Error()) - if werr := write("%-8s FAIL %v\n", l.Name, err); werr != nil { + if werr := fail(l.Name, err); werr != nil { return werr } continue } if err := preflightBinds([]bindTarget{{net: conf.Network, addr: conf.Addr}}); err != nil { - failures = append(failures, err.Error()) - if werr := write("%-8s FAIL %v\n", l.Name, err); werr != nil { + if werr := fail(l.Name, err); werr != nil { return werr } continue } + captureWanted = captureWanted || conf.Capture if err := write("%-8s OK %s %s %s\n", l.Name, conf.Network, conf.Addr, conf.HandlerSpec); err != nil { return err } } + if captureWanted { + if err := checkRunDir(); err != nil { + if werr := fail("capture", err); werr != nil { + return werr + } + } else if err := write("%-8s OK %s\n", "capture", "runs directory writable"); err != nil { + return err + } + } + if len(failures) > 0 { return fmt.Errorf("check failed:\n %s", strings.Join(failures, "\n ")) } @@ -201,7 +233,7 @@ func tryBind(network, addr string) error { func describeBindError(t bindTarget, err error) error { addr := t.addr if errors.Is(err, syscall.EACCES) || errors.Is(err, syscall.EPERM) { - if port, ok := parseAddrPortNumber(addr); ok && port < 1024 { + if port, ok := netx.ParsePort(addr); ok && port < 1024 { return fmt.Errorf("cannot bind %s: permission denied (ports below 1024 require elevated privileges on this system)", addr) } } @@ -211,18 +243,6 @@ func describeBindError(t bindTarget, err error) error { return fmt.Errorf("cannot bind %s: %w", addr, err) } -func parseAddrPortNumber(addr string) (int, bool) { - _, portStr, err := net.SplitHostPort(addr) - if err != nil { - return 0, false - } - port, err := strconv.Atoi(portStr) - if err != nil { - return 0, false - } - return port, true -} - func init() { rootCmd.AddCommand(checkCmd) } diff --git a/cmd/check_test.go b/cmd/check_test.go new file mode 100644 index 0000000..0142b25 --- /dev/null +++ b/cmd/check_test.go @@ -0,0 +1,28 @@ +package cmd + +import ( + "os" + "path/filepath" + "testing" + + "github.com/lachlanharrisdev/gonetsim/internal/capture" +) + +func TestCheckRunDir(t *testing.T) { + dir := t.TempDir() + t.Setenv("XDG_DATA_HOME", dir) + + if err := checkRunDir(); err != nil { + t.Fatalf("checkRunDir: %v", err) + } + runs, err := capture.DefaultRunsDir() + if err != nil { + t.Fatalf("DefaultRunsDir: %v", err) + } + if st, err := os.Stat(runs); err != nil || !st.IsDir() { + t.Fatalf("expected runs dir to exist: %v", err) + } + if filepath.Dir(runs) != filepath.Join(dir, "gonetsim") { + t.Fatalf("runs dir = %q, want it under %q", runs, dir) + } +} diff --git a/cmd/pcap.go b/cmd/pcap.go index ed592d2..face9f5 100644 --- a/cmd/pcap.go +++ b/cmd/pcap.go @@ -1,21 +1,104 @@ package cmd import ( + "fmt" + "io" + "io/fs" + "os" + "path/filepath" + "sort" + "strings" + "time" + "github.com/spf13/cobra" "github.com/lachlanharrisdev/gonetsim/internal/capture" ) var pcapCmd = &cobra.Command{ - Use: "pcap ", - Short: "Open a PCAP file", - Args: cobra.ExactArgs(1), - RunE: func(cmd *cobra.Command, args []string) error { - err := capture.ReadPcap(args[0]) - return err - }, + Use: "pcap ", + Short: "Inspect pcapng capture files", + Long: "Reads pcapng capture files and prints a summary. Pass a single file\n" + + "or a directory (e.g. the runs directory) to summarize every capture beneath it.\n" + + "Legacy pcap files are not supported", + Args: cobra.ExactArgs(1), + RunE: runPcap, } func init() { rootCmd.AddCommand(pcapCmd) } + +func runPcap(cmd *cobra.Command, args []string) error { + return inspectPcap(cmd.OutOrStdout(), args[0]) +} + +func inspectPcap(out io.Writer, target string) error { + st, err := os.Stat(target) + if err != nil { + return fmt.Errorf("pcap %q: %w", target, err) + } + if st.IsDir() { + return inspectPcapDir(out, target) + } + info, err := capture.Inspect(target) + if err != nil { + return err + } + fmt.Fprintf(out, "%s\n", summarizePcap(target, info)) + return nil +} + +func summarizePcap(path string, info capture.FileInfo) string { + var sb strings.Builder + fmt.Fprintf(&sb, "%s: format=pcapng linktype=%s packets=%d", path, info.LinkType, info.Packets) + if info.Packets > 0 { + fmt.Fprintf(&sb, " first=%s last=%s duration=%s", + info.First.Format(time.RFC3339), info.Last.Format(time.RFC3339), + info.Last.Sub(info.First).Round(time.Millisecond)) + } + if len(info.Interfaces) > 0 { + fmt.Fprintf(&sb, " interfaces=%s", strings.Join(info.Interfaces, "|")) + } + if info.CreatedBy != "" { + fmt.Fprintf(&sb, " app=%s", info.CreatedBy) + } + return sb.String() +} + +func inspectPcapDir(out io.Writer, dir string) error { + var files []string + walkErr := filepath.WalkDir(dir, func(path string, d fs.DirEntry, err error) error { + if err != nil { + return err + } + if !d.IsDir() && strings.HasSuffix(strings.ToLower(d.Name()), ".pcapng") { + files = append(files, path) + } + return nil + }) + if walkErr != nil { + return fmt.Errorf("pcap %q: %w", dir, walkErr) + } + sort.Strings(files) + if len(files) == 0 { + return fmt.Errorf("no pcapng files found in %q", dir) + } + var total uint64 + failed := 0 + for _, f := range files { + info, err := capture.Inspect(f) + if err != nil { + fmt.Fprintf(out, "%s: ERROR %v\n", f, err) + failed++ + continue + } + fmt.Fprintf(out, "%s\n", summarizePcap(f, info)) + total += info.Packets + } + fmt.Fprintf(out, "total: files=%d packets=%d\n", len(files), total) + if failed > 0 { + return fmt.Errorf("%d of %d files could not be read", failed, len(files)) + } + return nil +} diff --git a/cmd/pcap_test.go b/cmd/pcap_test.go new file mode 100644 index 0000000..37ab03d --- /dev/null +++ b/cmd/pcap_test.go @@ -0,0 +1,113 @@ +package cmd + +import ( + "bytes" + "net/netip" + "os" + "path/filepath" + "strings" + "testing" + + "github.com/lachlanharrisdev/gonetsim/internal/capture" +) + +func writePcapFixture(t *testing.T, path string, payloads ...string) { + t.Helper() + local := netip.MustParseAddrPort("127.0.0.1:8080") + remote := netip.MustParseAddrPort("10.0.0.5:40000") + run, err := capture.NewRun(path) + if err != nil { + t.Fatalf("NewRun: %v", err) + } + defer func() { _ = run.Close() }() + iface, err := run.NewInterface("test") + if err != nil { + t.Fatalf("NewInterface: %v", err) + } + ses, err := run.NewSession("tcp", local, remote, iface) + if err != nil { + t.Fatalf("NewSession: %v", err) + } + for _, p := range payloads { + if err := ses.Write([]byte(p), true); err != nil { + t.Fatalf("Write: %v", err) + } + } + if err := ses.Close(); err != nil { + t.Fatalf("Close: %v", err) + } +} + +func TestInspectPcapFile(t *testing.T) { + path := filepath.Join(t.TempDir(), "flow.pcapng") + writePcapFixture(t, path, "hello", "world") + + var out bytes.Buffer + if err := inspectPcap(&out, path); err != nil { + t.Fatalf("inspectPcap: %v", err) + } + got := out.String() + for _, want := range []string{"format=pcapng", "packets=", "first=", "last=", "duration="} { + if !strings.Contains(got, want) { + t.Errorf("output %q missing %q", got, want) + } + } +} + +func TestInspectPcapDir(t *testing.T) { + dir := t.TempDir() + writePcapFixture(t, filepath.Join(dir, "b.pcapng"), "one") + writePcapFixture(t, filepath.Join(dir, "a.pcapng"), "one", "two") + if err := os.WriteFile(filepath.Join(dir, "notes.txt"), []byte("ignore me"), 0o644); err != nil { + t.Fatalf("WriteFile: %v", err) + } + + var out bytes.Buffer + if err := inspectPcap(&out, dir); err != nil { + t.Fatalf("inspectPcap: %v", err) + } + got := out.String() + if !strings.Contains(got, "total: files=2") { + t.Errorf("missing totals line: %q", got) + } + if strings.Contains(got, "notes.txt") { + t.Errorf("non-pcapng file should be skipped: %q", got) + } + if a, b := strings.Index(got, "a.pcapng"), strings.Index(got, "b.pcapng"); a < 0 || b < 0 || a > b { + t.Errorf("files should be listed sorted: %q", got) + } +} + +func TestInspectPcapDirWithBadFile(t *testing.T) { + dir := t.TempDir() + writePcapFixture(t, filepath.Join(dir, "good.pcapng"), "one") + if err := os.WriteFile(filepath.Join(dir, "bad.pcapng"), []byte("not a capture"), 0o644); err != nil { + t.Fatalf("WriteFile: %v", err) + } + + var out bytes.Buffer + err := inspectPcap(&out, dir) + if err == nil || !strings.Contains(err.Error(), "could not be read") { + t.Fatalf("expected unreadable-file error, got %v", err) + } + if got := out.String(); !strings.Contains(got, "good.pcapng") || !strings.Contains(got, "bad.pcapng: ERROR") { + t.Errorf("good files should still be listed alongside errors: %q", got) + } +} + +func TestInspectPcapFailures(t *testing.T) { + var out bytes.Buffer + if err := inspectPcap(&out, filepath.Join(t.TempDir(), "empty")); err == nil { + t.Errorf("expected error for directory without captures") + } + if err := inspectPcap(&out, filepath.Join(t.TempDir(), "missing.pcapng")); err == nil { + t.Errorf("expected error for missing file") + } + bad := filepath.Join(t.TempDir(), "bad.pcapng") + if err := os.WriteFile(bad, []byte("not a capture"), 0o644); err != nil { + t.Fatalf("WriteFile: %v", err) + } + if err := inspectPcap(&out, bad); err == nil { + t.Errorf("expected error for corrupt file") + } +} diff --git a/cmd/run.go b/cmd/run.go index 911e9a2..bb5b735 100644 --- a/cmd/run.go +++ b/cmd/run.go @@ -13,6 +13,7 @@ import ( "github.com/spf13/cobra" + "github.com/lachlanharrisdev/gonetsim/internal/capture" appconfig "github.com/lachlanharrisdev/gonetsim/internal/config" "github.com/lachlanharrisdev/gonetsim/internal/observability" "github.com/lachlanharrisdev/gonetsim/internal/service" @@ -25,7 +26,7 @@ type runOptions struct { timeout time.Duration tls bool noCapture bool - artifacts string + output string } var runOpts runOptions @@ -40,9 +41,9 @@ func addRunFlags(cmd *cobra.Command) { cmd.Flags().BoolVar(&runOpts.tls, "tls", false, "wrap inline tcp listeners in TLS with an in-memory self-signed certificate") cmd.Flags().BoolVar(&runOpts.noCapture, "no-capture", false, - "disable capture for this run") - cmd.Flags().StringVar(&runOpts.artifacts, "artifacts", "", - "base directory for capture files (default ./artifacts)") + "don't write a capture file for this run") + cmd.Flags().StringVar(&runOpts.output, "output", "", + "write the run capture to this pcapng file instead of the default runs directory") } var runCmd = &cobra.Command{ @@ -85,7 +86,7 @@ func runTargets(cmd *cobra.Command, args []string) error { return err } - logger, err := observability.NewLogger(cfg.Logging) + logger, err := observability.NewLogger(observability.Options{Format: cfg.Logging.LogFormat, Level: cfg.Logging.Level}) if err != nil { return err } @@ -109,17 +110,36 @@ func runTargets(cmd *cobra.Command, args []string) error { return err } - limit, err := state.ParseSize(cfg.State.TotalLimit) + limit, err := appconfig.ParseSize(cfg.State.TotalLimit) if err != nil { return err } global := state.NewStore(state.NewBudget(limit)) - resolved, err := resolveTargets(specs, &cfg, configDir, cwd, runOpts, logger, global) + var run *capture.Run + if !runOpts.noCapture { + path, err := capture.RunPath(runOpts.output) + if err != nil { + return err + } + run, err = capture.NewRun(path) + if err != nil { + return err + } + logger.Info("capture", "path", path) + } + + resolved, err := resolveTargets(specs, &cfg, configDir, cwd, runOpts, logger, global, run) if err != nil { + if run != nil { + _ = run.Close() + } return err } if len(resolved) == 0 { + if run != nil { + _ = run.Close() + } return fmt.Errorf("at least one service must be enabled") } @@ -131,5 +151,12 @@ func runTargets(cmd *cobra.Command, args []string) error { } logger.Info("running", "targets", strings.Join(displays, " ")) - return manager.RunAll(runCtx) + runErr := manager.RunAll(runCtx) + if run != nil { + packets, first, last := run.Stats() + path := run.Path() + _ = run.Close() + logger.Info("capture saved", "path", path, "packets", packets, "duration", last.Sub(first).Round(time.Millisecond)) + } + return runErr } diff --git a/cmd/script.go b/cmd/script.go index 6380b6a..b087d07 100644 --- a/cmd/script.go +++ b/cmd/script.go @@ -23,7 +23,8 @@ var scriptCmd = &cobra.Command{ return err } - logger, err := observability.NewLogger(appconfig.Default().Logging) + def := appconfig.Default().Logging + logger, err := observability.NewLogger(observability.Options{Format: def.LogFormat, Level: def.Level}) if err != nil { return err } diff --git a/cmd/servicecfg.go b/cmd/servicecfg.go index b899997..c47e0c1 100644 --- a/cmd/servicecfg.go +++ b/cmd/servicecfg.go @@ -2,9 +2,7 @@ package cmd import ( "fmt" - "net" "net/netip" - "path/filepath" "strings" "time" @@ -12,21 +10,10 @@ import ( "github.com/lachlanharrisdev/gonetsim/internal/dnsserver" "github.com/lachlanharrisdev/gonetsim/internal/httpserver" "github.com/lachlanharrisdev/gonetsim/internal/listener" + "github.com/lachlanharrisdev/gonetsim/internal/netx" "github.com/lachlanharrisdev/gonetsim/internal/tlsprovider" ) -func parseAddrPort(listen string) (string, error) { - if listen == "" { - return "", fmt.Errorf("listen address is required") - } - - if _, err := net.ResolveTCPAddr("tcp", listen); err != nil { - return "", fmt.Errorf("invalid listen address %q (expected host:port): %w", listen, err) - } - - return listen, nil -} - func parseNetipAddr(s string) (netip.Addr, error) { a, err := netip.ParseAddr(s) if err != nil { @@ -45,7 +32,7 @@ func parseOptionalNetipAddr(s string) (netip.Addr, error) { const defaultReadTimeout = 30 * time.Second func listenerConfig(l appconfig.ListenerConfig, configDir string) (listener.Config, error) { - listen, err := parseAddrPort(l.Listen) + listen, err := netx.ParseAddr(l.Listen) if err != nil { return listener.Config{}, fmt.Errorf("listener %s.listen: %w", l.Name, err) } @@ -73,8 +60,7 @@ func listenerConfig(l appconfig.ListenerConfig, configDir string) (listener.Conf if l.TLS || l.TLSCert != "" || l.TLSKey != "" { certPath, keyPath := l.TLSCert, l.TLSKey if certPath == "" && keyPath == "" { - certPath = filepath.Join(configDir, tlsprovider.PersistedCertFileName) - keyPath = filepath.Join(configDir, tlsprovider.PersistedKeyFileName) + certPath, keyPath = tlsprovider.DefaultPaths(configDir) } conf.TLS = &tlsprovider.Config{CertFile: certPath, KeyFile: keyPath} } @@ -93,7 +79,7 @@ func dnsIPv4(s string) (netip.Addr, error) { } func dnsConfig(cfg appconfig.DNSConfig) (dnsserver.Config, error) { - listen, err := parseAddrPort(cfg.Listen) + listen, err := netx.ParseAddr(cfg.Listen) if err != nil { return dnsserver.Config{}, fmt.Errorf("dns.listen: %w", err) } @@ -114,6 +100,7 @@ func dnsConfig(cfg appconfig.DNSConfig) (dnsserver.Config, error) { SinkholeTXT: cfg.TXT, TTL: cfg.TTL, Compress: cfg.Compress, + Capture: cfg.Capture, } if err := conf.Validate(); err != nil { return dnsserver.Config{}, fmt.Errorf("dns: %w", err) @@ -122,7 +109,7 @@ func dnsConfig(cfg appconfig.DNSConfig) (dnsserver.Config, error) { } func httpConfig(cfg appconfig.HTTPConfig) (httpserver.Config, error) { - listen, err := parseAddrPort(cfg.Listen) + listen, err := netx.ParseAddr(cfg.Listen) if err != nil { return httpserver.Config{}, fmt.Errorf("http.listen: %w", err) } @@ -131,6 +118,7 @@ func httpConfig(cfg appconfig.HTTPConfig) (httpserver.Config, error) { StatusCode: cfg.Status, Mode: cfg.Mode, RootDir: cfg.RootDir, + Capture: cfg.Capture, } if err := conf.Validate(); err != nil { return httpserver.Config{}, fmt.Errorf("http: %w", err) @@ -139,21 +127,21 @@ func httpConfig(cfg appconfig.HTTPConfig) (httpserver.Config, error) { } func httpsConfig(cfg appconfig.HTTPSConfig, configDir string) (httpserver.Config, error) { - listen, err := parseAddrPort(cfg.Listen) + listen, err := netx.ParseAddr(cfg.Listen) if err != nil { return httpserver.Config{}, fmt.Errorf("https.listen: %w", err) } certPath := cfg.Cert keyPath := cfg.Key if certPath == "" && keyPath == "" { - certPath = filepath.Join(configDir, tlsprovider.PersistedCertFileName) - keyPath = filepath.Join(configDir, tlsprovider.PersistedKeyFileName) + certPath, keyPath = tlsprovider.DefaultPaths(configDir) } conf := httpserver.Config{ Addr: listen, StatusCode: cfg.Status, Mode: cfg.Mode, RootDir: cfg.RootDir, + Capture: cfg.Capture, TLS: &tlsprovider.Config{CertFile: certPath, KeyFile: keyPath}, } if err := conf.Validate(); err != nil { diff --git a/cmd/targets.go b/cmd/targets.go index 5ba903a..a98a420 100644 --- a/cmd/targets.go +++ b/cmd/targets.go @@ -3,15 +3,16 @@ package cmd import ( "fmt" "log/slog" - "net" "path/filepath" "strconv" "strings" + "github.com/lachlanharrisdev/gonetsim/internal/capture" appconfig "github.com/lachlanharrisdev/gonetsim/internal/config" "github.com/lachlanharrisdev/gonetsim/internal/dnsserver" "github.com/lachlanharrisdev/gonetsim/internal/httpserver" "github.com/lachlanharrisdev/gonetsim/internal/listener" + "github.com/lachlanharrisdev/gonetsim/internal/netx" "github.com/lachlanharrisdev/gonetsim/internal/service" "github.com/lachlanharrisdev/gonetsim/internal/state" "github.com/lachlanharrisdev/gonetsim/internal/tlsprovider" @@ -76,8 +77,8 @@ func parseInlineTarget(arg string) (targetSpec, error) { return targetSpec{}, fmt.Errorf("invalid network %q in %q (must be /tcp or /udp)", suffix, arg) } } - if _, err := net.ResolveTCPAddr("tcp", addr); err != nil { - return targetSpec{}, fmt.Errorf("invalid listen address %q in %q (expected host:port): %w", addr, arg, err) + if _, err := netx.ParseAddr(addr); err != nil { + return targetSpec{}, fmt.Errorf("invalid listen address in %q: %w", arg, err) } handlerSpec, name, err := resolveInlineHandler(spec) @@ -153,9 +154,9 @@ type resolvedTarget struct { display string } -func resolveTargets(specs []targetSpec, cfg *appconfig.Config, configDir, cwd string, opts runOptions, logger *slog.Logger, global *state.Store) ([]resolvedTarget, error) { +func resolveTargets(specs []targetSpec, cfg *appconfig.Config, configDir, cwd string, opts runOptions, logger *slog.Logger, global *state.Store, run *capture.Run) ([]resolvedTarget, error) { if len(specs) == 0 { - return resolveAll(cfg, configDir, opts, logger, global) + return resolveAll(cfg, configDir, opts, logger, global, run) } if opts.listen != "" && len(specs) > 1 { @@ -164,7 +165,7 @@ func resolveTargets(specs []targetSpec, cfg *appconfig.Config, configDir, cwd st out := make([]resolvedTarget, 0, len(specs)) for _, spec := range specs { - rt, err := resolveOne(spec, cfg, configDir, cwd, opts, logger, global) + rt, err := resolveOne(spec, cfg, configDir, cwd, opts, logger, global, run) if err != nil { return nil, err } @@ -173,13 +174,13 @@ func resolveTargets(specs []targetSpec, cfg *appconfig.Config, configDir, cwd st return out, nil } -func resolveAll(cfg *appconfig.Config, configDir string, opts runOptions, logger *slog.Logger, global *state.Store) ([]resolvedTarget, error) { +func resolveAll(cfg *appconfig.Config, configDir string, opts runOptions, logger *slog.Logger, global *state.Store, run *capture.Run) ([]resolvedTarget, error) { out := make([]resolvedTarget, 0, len(presetTargets)+len(cfg.Listeners)) for _, p := range presetTargets { if !p.enabled(cfg) { continue } - svc, display, err := p.build(cfg, configDir, opts, logger) + svc, display, err := p.build(cfg, configDir, opts, logger, run) if err != nil { return nil, err } @@ -190,7 +191,7 @@ func resolveAll(cfg *appconfig.Config, configDir string, opts runOptions, logger if !l.IsEnabled() { continue } - rt, err := resolveOne(targetSpec{raw: l.Name, kind: targetListener, name: l.Name}, cfg, configDir, "", opts, logger, global) + rt, err := resolveOne(targetSpec{raw: l.Name, kind: targetListener, name: l.Name}, cfg, configDir, "", opts, logger, global, run) if err != nil { return nil, err } @@ -199,14 +200,14 @@ func resolveAll(cfg *appconfig.Config, configDir string, opts runOptions, logger return out, nil } -func resolveOne(spec targetSpec, cfg *appconfig.Config, configDir, cwd string, opts runOptions, logger *slog.Logger, global *state.Store) (resolvedTarget, error) { +func resolveOne(spec targetSpec, cfg *appconfig.Config, configDir, cwd string, opts runOptions, logger *slog.Logger, global *state.Store, run *capture.Run) (resolvedTarget, error) { switch spec.kind { case targetPreset: for _, p := range presetTargets { if p.name != spec.preset { continue } - svc, display, err := p.build(cfg, configDir, opts, logger) + svc, display, err := p.build(cfg, configDir, opts, logger, run) if err != nil { return resolvedTarget{}, err } @@ -232,7 +233,7 @@ func resolveOne(spec targetSpec, cfg *appconfig.Config, configDir, cwd string, o if err := applyListenerRunOptions(&conf, opts, opts.listen); err != nil { return resolvedTarget{}, err } - return buildListener(conf, global, logger) + return buildListener(conf, global, logger, run) case targetInline: conf := spec.inline @@ -240,13 +241,19 @@ func resolveOne(spec targetSpec, cfg *appconfig.Config, configDir, cwd string, o if err := applyListenerRunOptions(&conf, opts, opts.listen); err != nil { return resolvedTarget{}, err } - return buildListener(conf, global, logger) + return buildListener(conf, global, logger, run) default: return resolvedTarget{}, fmt.Errorf("unknown target kind %d", spec.kind) } } +func applyCaptureOptions(noCapture bool, capture *bool) { + if noCapture { + *capture = false + } +} + func applyListenerRunOptions(conf *listener.Config, opts runOptions, listen string) error { if listen != "" { conf.Addr = listen @@ -260,9 +267,6 @@ func applyListenerRunOptions(conf *listener.Config, opts runOptions, listen stri if opts.noCapture { conf.Capture = false } - if opts.artifacts != "" { - conf.CaptureDir = opts.artifacts - } if opts.tls { if conf.Network != "tcp" { return fmt.Errorf("listener %s: --tls requires a tcp listener", conf.Name) @@ -273,8 +277,8 @@ func applyListenerRunOptions(conf *listener.Config, opts runOptions, listen stri return nil } -func buildListener(conf listener.Config, global *state.Store, logger *slog.Logger) (resolvedTarget, error) { - svc, err := listener.NewService(conf, global, logger) +func buildListener(conf listener.Config, global *state.Store, logger *slog.Logger, run *capture.Run) (resolvedTarget, error) { + svc, err := listener.NewService(conf, global, logger, run) if err != nil { return resolvedTarget{}, err } @@ -292,66 +296,75 @@ func listenerDisplay(conf listener.Config) string { return display + ")" } +func presetBuild[AC any, SC any]( + appCfg AC, + opts runOptions, + configDir string, + logger *slog.Logger, + run *capture.Run, + setListen func(*AC, string), + parse func(AC, string) (SC, error), + applyCapture func(*SC), + svc func(SC, *slog.Logger, *capture.Run) service.Service, + display func(SC) string, +) (service.Service, string, error) { + if opts.listen != "" { + setListen(&appCfg, opts.listen) + } + conf, err := parse(appCfg, configDir) + if err != nil { + return nil, "", err + } + applyCapture(&conf) + return svc(conf, logger, run), display(conf), nil +} + var presetTargets = []struct { name string enabled func(c *appconfig.Config) bool - build func(c *appconfig.Config, configDir string, opts runOptions, logger *slog.Logger) (service.Service, string, error) + build func(c *appconfig.Config, configDir string, opts runOptions, logger *slog.Logger, run *capture.Run) (service.Service, string, error) }{ { name: "dns", enabled: func(c *appconfig.Config) bool { return c.DNS.Enabled }, - build: func(c *appconfig.Config, _ string, opts runOptions, logger *slog.Logger) (service.Service, string, error) { - if opts.listen != "" { - c.DNS.Listen = opts.listen - } - conf, err := dnsConfig(c.DNS) - if err != nil { - return nil, "", err - } - return dnsserver.NewService(conf, logger), fmt.Sprintf("dns(%s/%s)", conf.Addr, netLabel(conf.Net)), nil + build: func(c *appconfig.Config, configDir string, opts runOptions, logger *slog.Logger, run *capture.Run) (service.Service, string, error) { + return presetBuild(c.DNS, opts, configDir, logger, run, + func(a *appconfig.DNSConfig, l string) { a.Listen = l }, + func(a appconfig.DNSConfig, _ string) (dnsserver.Config, error) { return dnsConfig(a) }, + func(s *dnsserver.Config) { applyCaptureOptions(opts.noCapture, &s.Capture) }, + dnsserver.NewService, + func(s dnsserver.Config) string { return fmt.Sprintf("dns(%s/%s)", s.Addr, netx.DisplayNetwork(s.Net)) }, + ) }, }, { name: "http", enabled: func(c *appconfig.Config) bool { return c.HTTP.Enabled }, - build: func(c *appconfig.Config, _ string, opts runOptions, logger *slog.Logger) (service.Service, string, error) { - if opts.listen != "" { - c.HTTP.Listen = opts.listen - } - conf, err := httpConfig(c.HTTP) - if err != nil { - return nil, "", err - } - return httpserver.NewService(conf, logger), fmt.Sprintf("http(%s)", conf.Addr), nil + build: func(c *appconfig.Config, configDir string, opts runOptions, logger *slog.Logger, run *capture.Run) (service.Service, string, error) { + return presetBuild(c.HTTP, opts, configDir, logger, run, + func(a *appconfig.HTTPConfig, l string) { a.Listen = l }, + func(a appconfig.HTTPConfig, _ string) (httpserver.Config, error) { return httpConfig(a) }, + func(s *httpserver.Config) { applyCaptureOptions(opts.noCapture, &s.Capture) }, + httpserver.NewService, + func(s httpserver.Config) string { return fmt.Sprintf("http(%s)", s.Addr) }, + ) }, }, { name: "https", enabled: func(c *appconfig.Config) bool { return c.HTTPS.Enabled }, - build: func(c *appconfig.Config, configDir string, opts runOptions, logger *slog.Logger) (service.Service, string, error) { - if opts.listen != "" { - c.HTTPS.Listen = opts.listen - } - conf, err := httpsConfig(c.HTTPS, configDir) - if err != nil { - return nil, "", err - } - return httpserver.NewService(conf, logger), fmt.Sprintf("https(%s)", conf.Addr), nil + build: func(c *appconfig.Config, configDir string, opts runOptions, logger *slog.Logger, run *capture.Run) (service.Service, string, error) { + return presetBuild(c.HTTPS, opts, configDir, logger, run, + func(a *appconfig.HTTPSConfig, l string) { a.Listen = l }, + func(a appconfig.HTTPSConfig, dir string) (httpserver.Config, error) { return httpsConfig(a, dir) }, + func(s *httpserver.Config) { applyCaptureOptions(opts.noCapture, &s.Capture) }, + httpserver.NewService, + func(s httpserver.Config) string { return fmt.Sprintf("https(%s)", s.Addr) }, + ) }, }, } -func netLabel(net string) string { - switch strings.ToLower(strings.TrimSpace(net)) { - case "both": - return "udp+tcp" - case "tcp": - return "tcp" - default: - return "udp" - } -} - func availableTargets(cfg *appconfig.Config) string { names := append([]string{}, presetNames...) for _, l := range cfg.Listeners { diff --git a/cmd/targets_test.go b/cmd/targets_test.go index 9e4d5f8..4183ce8 100644 --- a/cmd/targets_test.go +++ b/cmd/targets_test.go @@ -1,17 +1,17 @@ package cmd import ( - "io" "log/slog" "strings" "testing" appconfig "github.com/lachlanharrisdev/gonetsim/internal/config" "github.com/lachlanharrisdev/gonetsim/internal/state" + "github.com/lachlanharrisdev/gonetsim/internal/testutil" ) func testLogger() *slog.Logger { - return slog.New(slog.NewTextHandler(io.Discard, nil)) + return testutil.Logger() } func disabledAll(cfg *appconfig.Config) { @@ -104,10 +104,7 @@ func TestParseSets(t *testing.T) { func TestResolveTargets(t *testing.T) { t.Run("all enabled", func(t *testing.T) { cfg := appconfig.Default() - resolved, err := resolveTargets(nil, &cfg, t.TempDir(), t.TempDir(), runOptions{}, testLogger(), state.NewStore(nil)) - if err != nil { - t.Fatalf("resolveTargets: %v", err) - } + resolved := testResolve(t, &cfg, nil, runOptions{}) if len(resolved) != len(presetNames) { t.Fatalf("expected %d presets, got %d", len(presetNames), len(resolved)) } @@ -122,23 +119,20 @@ func TestResolveTargets(t *testing.T) { {Name: "off", Enabled: &disabled, Type: "tcp", Listen: "127.0.0.1:0", Handler: "builtin:sink"}, } - resolved, err := resolveTargets(nil, &cfg, t.TempDir(), t.TempDir(), runOptions{}, testLogger(), state.NewStore(nil)) - if err != nil { - t.Fatalf("resolveTargets: %v", err) - } + resolved := testResolve(t, &cfg, nil, runOptions{}) if len(resolved) != 1 { t.Fatalf("expected disabled listener to be skipped, got %d targets", len(resolved)) } - resolved, err = resolveTargets(mustSpecs(t, "off"), &cfg, t.TempDir(), t.TempDir(), runOptions{}, testLogger(), state.NewStore(nil)) - if err != nil || len(resolved) != 1 { - t.Fatalf("explicit target: %v, %d targets", err, len(resolved)) + resolved = testResolve(t, &cfg, []string{"off"}, runOptions{}) + if len(resolved) != 1 { + t.Fatalf("explicit target: %d targets", len(resolved)) } }) t.Run("unknown target lists alternatives", func(t *testing.T) { cfg := appconfig.Default() - _, err := resolveTargets(mustSpecs(t, "nope"), &cfg, t.TempDir(), t.TempDir(), runOptions{}, testLogger(), state.NewStore(nil)) + _, err := resolveTargets(mustSpecs(t, "nope"), &cfg, t.TempDir(), t.TempDir(), runOptions{}, testLogger(), state.NewStore(nil), nil) if err == nil || !strings.Contains(err.Error(), "unknown target") || !strings.Contains(err.Error(), "dns") || !strings.Contains(err.Error(), "handler@addr") { t.Fatalf("expected helpful unknown-target error, got: %v", err) @@ -147,7 +141,7 @@ func TestResolveTargets(t *testing.T) { t.Run("inline lua resolves script", func(t *testing.T) { cfg := appconfig.Default() - resolved, err := resolveTargets(mustSpecs(t, "testdata/hello.lua@127.0.0.1:0"), &cfg, t.TempDir(), ".", runOptions{}, testLogger(), state.NewStore(nil)) + resolved, err := resolveTargets(mustSpecs(t, "testdata/hello.lua@127.0.0.1:0"), &cfg, t.TempDir(), ".", runOptions{}, testLogger(), state.NewStore(nil), nil) if err != nil || len(resolved) != 1 || resolved[0].display != "hello(127.0.0.1:0)" { t.Fatalf("inline lua: %v, %+v", err, resolved) } @@ -155,7 +149,7 @@ func TestResolveTargets(t *testing.T) { t.Run("tls on udp rejected", func(t *testing.T) { cfg := appconfig.Default() - _, err := resolveTargets(mustSpecs(t, "sink@:0/udp"), &cfg, t.TempDir(), t.TempDir(), runOptions{tls: true}, testLogger(), state.NewStore(nil)) + _, err := resolveTargets(mustSpecs(t, "sink@:0/udp"), &cfg, t.TempDir(), t.TempDir(), runOptions{tls: true}, testLogger(), state.NewStore(nil), nil) if err == nil { t.Fatalf("expected --tls on udp to be rejected") } @@ -164,7 +158,7 @@ func TestResolveTargets(t *testing.T) { t.Run("listen requires single target", func(t *testing.T) { cfg := appconfig.Default() opts := runOptions{listen: "127.0.0.1:1234"} - _, err := resolveTargets(mustSpecs(t, "http", "dns"), &cfg, t.TempDir(), t.TempDir(), opts, testLogger(), state.NewStore(nil)) + _, err := resolveTargets(mustSpecs(t, "http", "dns"), &cfg, t.TempDir(), t.TempDir(), opts, testLogger(), state.NewStore(nil), nil) if err == nil { t.Fatalf("expected --listen with multiple targets to fail") } @@ -199,11 +193,11 @@ func TestServiceConfigMapping(t *testing.T) { t.Run("dns auto ipv4", func(t *testing.T) { cfg := appconfig.DNSConfig{ - Listen: "127.0.0.1:0", - Network: "udp", - IPv4: "auto", - Domain: "localhost", - TXT: "test", + ServiceBase: appconfig.ServiceBase{Listen: "127.0.0.1:0"}, + Network: "udp", + IPv4: "auto", + Domain: "localhost", + TXT: "test", } conf, err := dnsConfig(cfg) if err != nil { @@ -236,3 +230,16 @@ func mustSpecs(t *testing.T, args ...string) []targetSpec { } return specs } + +func testResolve(t *testing.T, cfg *appconfig.Config, args []string, opts runOptions) []resolvedTarget { + t.Helper() + var specs []targetSpec + if len(args) > 0 { + specs = mustSpecs(t, args...) + } + resolved, err := resolveTargets(specs, cfg, t.TempDir(), t.TempDir(), opts, testLogger(), state.NewStore(nil), nil) + if err != nil { + t.Fatalf("resolveTargets(%v): %v", args, err) + } + return resolved +} diff --git a/docker/docker-compose.yml b/docker/docker-compose.yml index 2cb89b4..8cdb47b 100644 --- a/docker/docker-compose.yml +++ b/docker/docker-compose.yml @@ -32,7 +32,7 @@ services: # - ../gonetsim.toml:/etc/gonetsim/gonetsim.toml:ro # - ../handlers:/etc/gonetsim/handlers:ro - # listener capture files are written to ./artifacts inside the container - # (i.e. /artifacts); mount a volume there to keep them: + # the run capture is written inside the container's data directory; + # mount a volume there to keep it, or pass --output to choose a path: # volumes: - # - ./artifacts:/artifacts + # - ./captures:/root/.local/share/gonetsim/runs diff --git a/examples/handlers/ftp.lua b/examples/handlers/ftp.lua index 68c1c37..665b329 100644 --- a/examples/handlers/ftp.lua +++ b/examples/handlers/ftp.lua @@ -22,7 +22,7 @@ function handle(conn) line = line:gsub("%s+$", "") if line ~= "" then - capture:write("ftp", line) + capture:comment("ftp: " .. line) local cmd = line:match("^(%S+)") local arg = line:match("^%S+%s+(.+)$") diff --git a/examples/handlers/irc.lua b/examples/handlers/irc.lua index 9c8d456..c3ec907 100644 --- a/examples/handlers/irc.lua +++ b/examples/handlers/irc.lua @@ -22,7 +22,7 @@ function handle(conn) line = line:gsub("%s+$", "") if line ~= "" then - capture:write("irc", line) + capture:comment("irc: " .. line) log:info(line) local cmd = line:match("^(%S+)") diff --git a/examples/handlers/smtp.lua b/examples/handlers/smtp.lua index bd14e39..3d34797 100644 --- a/examples/handlers/smtp.lua +++ b/examples/handlers/smtp.lua @@ -100,7 +100,7 @@ local function doAuth(conn, arg) -- PLAIN payload is authzid NUL authcid NUL passwd user, pass = parts[2] or "?", parts[3] or "?" end - capture:write("auth", user .. " / " .. pass) + capture:comment("auth: " .. user .. " / " .. pass) log:info("AUTH " .. mech .. " captured") conn:write("235 2.7.0 Authentication successful\r\n") end @@ -116,7 +116,7 @@ function handle(conn) line = line:gsub("%s+$", "") if line ~= "" then - capture:write("smtp", line) + capture:comment("smtp: " .. line) end local cmd = line:match("^(%a+)") or "" @@ -151,7 +151,7 @@ function handle(conn) if not msg then break end msg = msg:gsub("\r\n%.\r\n$", "\r\n") -- strip the terminator msg = msg:gsub("\r\n%.%.", "\r\n.") -- un-dot-stuff - capture:write("message", msg) + capture:comment("message: " .. msg) local mails = tonumber(handler:get("mails")) or 0 mails = mails + 1 handler:set("mails", tostring(mails)) diff --git a/go.mod b/go.mod index 3228721..e5e59d1 100644 --- a/go.mod +++ b/go.mod @@ -4,6 +4,7 @@ go 1.26.1 require ( github.com/fatih/color v1.19.0 + github.com/google/gopacket v1.1.19 github.com/knadh/koanf/parsers/toml/v2 v2.2.2 github.com/knadh/koanf/providers/confmap v1.0.1 github.com/knadh/koanf/providers/file v1.2.1 @@ -19,7 +20,6 @@ require ( require ( github.com/fsnotify/fsnotify v1.10.1 // indirect github.com/go-viper/mapstructure/v2 v2.5.0 // indirect - github.com/google/gopacket v1.1.19 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/knadh/koanf/maps v0.1.3 // indirect github.com/mitchellh/copystructure v1.2.0 // indirect diff --git a/internal/capture/capture.go b/internal/capture/capture.go deleted file mode 100644 index 796de31..0000000 --- a/internal/capture/capture.go +++ /dev/null @@ -1,73 +0,0 @@ -package capture - -import ( - "fmt" - "os" - "path/filepath" - "strings" - "time" -) - -// base directory capture files are written to, relative to cwd -const DefaultDir = "artifacts" - -type Store struct { - dir string -} - -func NewStore(baseDir, listener string) (*Store, error) { - dir := filepath.Join(baseDir, sanitize(listener)) - if err := os.MkdirAll(dir, 0o755); err != nil { - return nil, fmt.Errorf("create capture dir %q: %w", dir, err) - } - return &Store{dir: dir}, nil -} - -func (s *Store) Conn(remote string, now time.Time) (*Writer, error) { - if s == nil { - return nil, nil - } - name := now.Format("20060102-150405.000000") + "-" + sanitize(remote) + ".log" - f, err := os.OpenFile(filepath.Join(s.dir, name), os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0o644) - if err != nil { - return nil, fmt.Errorf("create capture file: %w", err) - } - return &Writer{f: f}, nil -} - -type Writer struct { - f *os.File -} - -func (w *Writer) Write(name string, data []byte) { - if w == nil { - return - } - if name != "" { - _, _ = fmt.Fprintf(w.f, "=== %s ===\n", name) - } - _, _ = w.f.Write(data) - if name != "" && (len(data) == 0 || data[len(data)-1] != '\n') { - _, _ = w.f.Write([]byte{'\n'}) - } -} - -func (w *Writer) Close() error { - if w == nil || w.f == nil { - return nil - } - err := w.f.Close() - w.f = nil - return err -} - -func sanitize(s string) string { - return strings.Map(func(r rune) rune { - switch { - case r >= 'a' && r <= 'z', r >= 'A' && r <= 'Z', r >= '0' && r <= '9', r == '.', r == '_', r == '-': - return r - default: - return '_' - } - }, s) -} diff --git a/internal/capture/capture_test.go b/internal/capture/capture_test.go deleted file mode 100644 index 1a2c4c4..0000000 --- a/internal/capture/capture_test.go +++ /dev/null @@ -1,82 +0,0 @@ -package capture - -import ( - "os" - "path/filepath" - "strings" - "testing" - "time" -) - -func TestCapture(t *testing.T) { - t.Run("connection files", func(t *testing.T) { - base := t.TempDir() - store, err := NewStore(base, "test") - if err != nil { - t.Fatalf("NewStore: %v", err) - } - - w, err := store.Conn("203.0.113.10:43210", time.Now()) - if err != nil { - t.Fatalf("Conn: %v", err) - } - w.Write("", []byte("captured data")) - if err := w.Close(); err != nil { - t.Fatalf("Close: %v", err) - } - - entries, err := os.ReadDir(filepath.Join(base, "test")) - if err != nil || len(entries) != 1 { - t.Fatalf("expected 1 capture file, got %d (%v)", len(entries), err) - } - if !strings.HasSuffix(entries[0].Name(), "203.0.113.10_43210.log") { - t.Fatalf("unexpected capture file name %q", entries[0].Name()) - } - data, _ := os.ReadFile(filepath.Join(base, "test", entries[0].Name())) - if string(data) != "captured data" { - t.Fatalf("unexpected capture content %q", data) - } - }) - - t.Run("named sections", func(t *testing.T) { - path := filepath.Join(t.TempDir(), "capture.log") - f, err := os.Create(path) - if err != nil { - t.Fatalf("Create: %v", err) - } - w := &Writer{f: f} - w.Write("request", []byte("GET /")) - w.Write("request", []byte("multi\nline\n")) - w.Write("", []byte("raw")) - if err := f.Close(); err != nil { - t.Fatalf("Close: %v", err) - } - - data, err := os.ReadFile(path) - if err != nil { - t.Fatalf("ReadFile: %v", err) - } - want := "=== request ===\nGET /\n=== request ===\nmulti\nline\nraw" - if string(data) != want { - t.Fatalf("unexpected content:\n%q\nwant:\n%q", data, want) - } - }) - - t.Run("nil values are no-ops", func(t *testing.T) { - var store *Store - w, err := store.Conn("1.2.3.4:5", time.Now()) - if err != nil || w != nil { - t.Fatalf("expected nil writer from nil store, got %v, %v", w, err) - } - w.Write("name", []byte("data")) - if err := w.Close(); err != nil { - t.Fatalf("Close on nil writer: %v", err) - } - }) - - t.Run("sanitize", func(t *testing.T) { - if got := sanitize("2001:db8::1%eth0/a b"); got != "2001_db8__1_eth0_a_b" { - t.Fatalf("unexpected sanitized name %q", got) - } - }) -} diff --git a/internal/capture/inspect.go b/internal/capture/inspect.go new file mode 100644 index 0000000..0eda1a2 --- /dev/null +++ b/internal/capture/inspect.go @@ -0,0 +1,100 @@ +package capture + +import ( + "encoding/binary" + "errors" + "fmt" + "io" + "os" + "strings" + "time" + + "github.com/google/gopacket/layers" + "github.com/google/gopacket/pcapgo" +) + +type FileInfo struct { + LinkType layers.LinkType + Packets uint64 + First time.Time + Last time.Time + CreatedBy string + Interfaces []string +} + +// true if b holds the magic bytes of a legacy pcap +// for either byte order and either timestamp resolution +func isLegacyMagic(b []byte) bool { + return b[0] == 0xa1 && b[1] == 0xb2 && b[2] == 0xc3 && b[3] == 0xd4 || + b[0] == 0xa1 && b[1] == 0xb2 && b[2] == 0x3c && b[3] == 0x4d || + b[3] == 0xa1 && b[2] == 0xb2 && b[1] == 0xc3 && b[0] == 0xd4 || + b[3] == 0xa1 && b[2] == 0xb2 && b[1] == 0x3c && b[0] == 0x4d +} + +func isHeaderOnly(f *os.File) (bool, error) { + st, err := f.Stat() + if err != nil { + return false, err + } + var hdr [8]byte + if _, err := f.ReadAt(hdr[:], 0); err != nil { + return false, err + } + blockType := binary.LittleEndian.Uint32(hdr[0:4]) + blockLen := int64(binary.LittleEndian.Uint32(hdr[4:8])) + return blockType == 0x0A0D0D0A && blockLen == st.Size(), nil +} + +func Inspect(path string) (FileInfo, error) { + f, err := os.Open(path) + if err != nil { + return FileInfo{}, fmt.Errorf("open %q: %w", path, err) + } + defer f.Close() + + var magic [4]byte + if _, err := io.ReadFull(f, magic[:]); err != nil { + return FileInfo{}, fmt.Errorf("%q is not a pcapng file: %w", path, err) + } + if isLegacyMagic(magic[:]) { + return FileInfo{}, fmt.Errorf("%q is a legacy pcap file; pcapng is the only supported format", path) + } + if magic != [4]byte{0x0a, 0x0d, 0x0d, 0x0a} { + return FileInfo{}, fmt.Errorf("%q is not a pcapng file", path) + } + + if _, err := f.Seek(0, io.SeekStart); err != nil { + return FileInfo{}, fmt.Errorf("seek %q: %w", path, err) + } + nr, err := pcapgo.NewNgReader(f, pcapgo.DefaultNgReaderOptions) + if err != nil { + if empty, serr := isHeaderOnly(f); serr == nil && empty { + return FileInfo{LinkType: layers.LinkTypeEthernet}, nil + } + return FileInfo{}, fmt.Errorf("read pcapng %q: %w", path, err) + } + + info := FileInfo{LinkType: nr.LinkType(), CreatedBy: nr.SectionInfo().Application} + for { + _, ci, err := nr.ReadPacketData() + if errors.Is(err, io.EOF) { + break + } + if err != nil { + return FileInfo{}, fmt.Errorf("read packet %d from %q: %w", info.Packets+1, path, err) + } + if info.Packets == 0 { + info.First = ci.Timestamp + } + info.Last = ci.Timestamp + info.Packets++ + } + for i := 0; i < nr.NInterfaces(); i++ { + iface, err := nr.Interface(i) + if err != nil { + break + } + info.Interfaces = append(info.Interfaces, strings.TrimRight(iface.Name, "\x00")) + } + return info, nil +} diff --git a/internal/capture/pcapng.go b/internal/capture/pcapng.go deleted file mode 100644 index a197a63..0000000 --- a/internal/capture/pcapng.go +++ /dev/null @@ -1,71 +0,0 @@ -package capture - -import ( - "encoding/binary" - "fmt" - "log" - "os" - - "github.com/google/gopacket" - "github.com/google/gopacket/pcap" - "github.com/google/gopacket/pcapgo" -) - -func ReadPcap(pcapFile string) error { - f, err := os.Open(pcapFile) - if err != nil { - return err - } - defer f.Close() - - buf := make([]byte, 4) - _, err = f.ReadAt(buf, 0) - if err != nil { - return err - } - - if binary.BigEndian.Uint32(buf) == 0x0a0d0d0a { - // pcapng - return ReadPcapNG(f) - } - magic := binary.BigEndian.Uint32(buf) - littleMagic := binary.LittleEndian.Uint32(buf) - - if magic == 0xa1b2c3d4 || magic == 0xa1b23c4d || - littleMagic == 0xa1b2c3d4 || littleMagic == 0xa1b23c4d { - //pcap - return ReadLegacyPcap(pcapFile) - } - - return fmt.Errorf("unknown pcap header") -} - -func ReadPcapNG(f *os.File) error { - reader, err := pcapgo.NewNgReader(f, pcapgo.DefaultNgReaderOptions) - if err != nil { - return err - } - - packetSource := gopacket.NewPacketSource(reader, reader.LinkType()) - for packet := range packetSource.Packets() { - fmt.Println(packet) - } - - return nil -} - -func ReadLegacyPcap(f string) error { - handle, err := pcap.OpenOffline(f) - if err != nil { - log.Fatal(err) - } - defer handle.Close() - - // Loop through packets in file - packetSource := gopacket.NewPacketSource(handle, handle.LinkType()) - for packet := range packetSource.Packets() { - fmt.Println(packet) - } - - return nil -} diff --git a/internal/capture/recorder.go b/internal/capture/recorder.go new file mode 100644 index 0000000..41f462e --- /dev/null +++ b/internal/capture/recorder.go @@ -0,0 +1,203 @@ +package capture + +import ( + "context" + "crypto/tls" + "net" + "net/netip" + "sync" + "time" +) + +type Conn struct { + net.Conn + ses *Session +} + +type ConnListener struct { + net.Listener + run *Run + iface int +} + +type udpFlow struct { + ses *Session + last time.Time +} + +type PacketConn struct { + net.PacketConn + run *Run + iface int + idle time.Duration + + mu sync.Mutex + flows map[string]*udpFlow +} + +func NewConnListener(ln net.Listener, run *Run, iface int) net.Listener { + if run == nil { + return ln + } + return &ConnListener{Listener: ln, run: run, iface: iface} +} + +func (l *ConnListener) Accept() (net.Conn, error) { + c, err := l.Listener.Accept() + if err != nil { + return nil, err + } + return NewConn(c, l.run, l.iface), nil +} + +func NewConn(c net.Conn, run *Run, iface int) net.Conn { + if run == nil { + return c + } + local, okL := toAddrPort(c.LocalAddr()) + remote, okR := toAddrPort(c.RemoteAddr()) + if !okL || !okR { + return c + } + ses, err := run.NewSession("tcp", local, remote, iface) + if err != nil || ses == nil { + return c + } + return &Conn{Conn: c, ses: ses} +} + +func (c *Conn) Session() *Session { + if c == nil { + return nil + } + return c.ses +} + +func (c *Conn) Read(p []byte) (int, error) { + n, err := c.Conn.Read(p) + if n > 0 { + _ = c.ses.Write(p[:n], true) + } + return n, err +} + +func (c *Conn) Write(p []byte) (int, error) { + n, err := c.Conn.Write(p) + if n > 0 { + _ = c.ses.Write(p[:n], false) + } + return n, err +} + +func (c *Conn) Close() error { + err := c.Conn.Close() + if c.ses != nil { + _ = c.ses.Close() + } + return err +} + +func (c *Conn) ConnectionState() tls.ConnectionState { + if tc, ok := c.Conn.(interface{ ConnectionState() tls.ConnectionState }); ok { + return tc.ConnectionState() + } + return tls.ConnectionState{} +} + +func (c *Conn) HandshakeContext(ctx context.Context) error { + if tc, ok := c.Conn.(interface { + HandshakeContext(ctx context.Context) error + }); ok { + return tc.HandshakeContext(ctx) + } + return nil +} + +func NewPacketConn(pc net.PacketConn, run *Run, iface int, idle time.Duration) *PacketConn { + return &PacketConn{PacketConn: pc, run: run, iface: iface, idle: idle, flows: make(map[string]*udpFlow)} +} + +func (c *PacketConn) SessionFor(remote net.Addr) *Session { + if c == nil { + return nil + } + c.mu.Lock() + defer c.mu.Unlock() + if f, ok := c.flows[remote.String()]; ok { + return f.ses + } + return nil +} + +func (c *PacketConn) ReadFrom(p []byte) (int, net.Addr, error) { + n, addr, err := c.PacketConn.ReadFrom(p) + if n > 0 { + c.record(addr, p[:n], true) + } + return n, addr, err +} + +func (c *PacketConn) WriteTo(p []byte, addr net.Addr) (int, error) { + n, err := c.PacketConn.WriteTo(p, addr) + if n > 0 { + c.record(addr, p[:n], false) + } + return n, err +} + +func (c *PacketConn) record(remote net.Addr, data []byte, fromClient bool) { + if c.run == nil { + return + } + c.mu.Lock() + defer c.mu.Unlock() + + now := time.Now() + if c.idle > 0 { + for k, f := range c.flows { + if now.Sub(f.last) > c.idle { + _ = f.ses.Close() + delete(c.flows, k) + } + } + } + + key := remote.String() + f, ok := c.flows[key] + if !ok { + local, okL := toAddrPort(c.LocalAddr()) + rem, okR := toAddrPort(remote) + if !okL || !okR { + return + } + ses, err := c.run.NewSession("udp", local, rem, c.iface) + if err != nil || ses == nil { + return + } + f = &udpFlow{ses: ses} + c.flows[key] = f + } + f.last = now + _ = f.ses.Write(data, fromClient) + _ = f.ses.Flush() +} + +func (c *PacketConn) CloseAll() { + c.mu.Lock() + defer c.mu.Unlock() + for k, f := range c.flows { + _ = f.ses.Close() + delete(c.flows, k) + } +} + +func toAddrPort(a net.Addr) (netip.AddrPort, bool) { + switch v := a.(type) { + case *net.TCPAddr: + return v.AddrPort(), true + case *net.UDPAddr: + return v.AddrPort(), true + default: + return netip.AddrPort{}, false + } +} diff --git a/internal/capture/run.go b/internal/capture/run.go new file mode 100644 index 0000000..9e4d673 --- /dev/null +++ b/internal/capture/run.go @@ -0,0 +1,188 @@ +package capture + +import ( + "crypto/rand" + "encoding/binary" + "fmt" + "net/netip" + "os" + "path/filepath" + "runtime" + "sync" + "time" + + "github.com/google/gopacket/layers" +) + +type Run struct { + mu sync.Mutex + f *os.File + path string + ifaces int + packets uint64 + first time.Time + last time.Time +} + +func NewRunID() string { + var suffix [2]byte + _, _ = rand.Read(suffix[:]) + return time.Now().Format("20060102-150405") + fmt.Sprintf("-%02x%02x", suffix[0], suffix[1]) +} + +func DefaultRunsDir() (string, error) { + var base string + switch runtime.GOOS { + case "windows": + dir, err := os.UserCacheDir() + if err != nil { + return "", err + } + base = dir + case "darwin": + home, err := os.UserHomeDir() + if err != nil { + return "", err + } + base = filepath.Join(home, "Library", "Application Support") + default: + if xdg := os.Getenv("XDG_DATA_HOME"); xdg != "" { + base = xdg + } else { + home, err := os.UserHomeDir() + if err != nil { + return "", err + } + base = filepath.Join(home, ".local", "share") + } + } + return filepath.Join(base, "gonetsim", "runs"), nil +} + +func RunPath(output string) (string, error) { + if output != "" { + if dir := filepath.Dir(output); dir != "." && dir != "" { + if err := os.MkdirAll(dir, 0o755); err != nil { + return "", fmt.Errorf("create output dir %q: %w", dir, err) + } + } + return output, nil + } + dir, err := DefaultRunsDir() + if err != nil { + return "", err + } + if err := os.MkdirAll(dir, 0o755); err != nil { + return "", fmt.Errorf("create runs dir %q: %w", dir, err) + } + return filepath.Join(dir, NewRunID()+".pcapng"), nil +} + +func NewRun(path string) (*Run, error) { + f, err := os.Create(path) + if err != nil { + return nil, fmt.Errorf("create pcapng %q: %w", path, err) + } + r := &Run{f: f, path: path} + if err := r.writeSHB(); err != nil { + _ = f.Close() + _ = os.Remove(path) + return nil, err + } + return r, nil +} + +func (r *Run) Path() string { + if r == nil { + return "" + } + return r.path +} + +func (r *Run) NewInterface(name string) (int, error) { + if r == nil { + return 0, nil + } + r.mu.Lock() + defer r.mu.Unlock() + opt := encodeOption(2, append([]byte(name), 0)) + opt = append(opt, encodeOption(0, nil)...) + length := 16 + len(opt) + 4 + b := make([]byte, 16) + binary.LittleEndian.PutUint32(b[0:4], 1) + binary.LittleEndian.PutUint32(b[4:8], uint32(length)) + binary.LittleEndian.PutUint16(b[8:10], uint16(layers.LinkTypeEthernet)) + binary.LittleEndian.PutUint16(b[10:12], 0) + binary.LittleEndian.PutUint32(b[12:16], snapLen) + if err := writeAll(r.f, b); err != nil { + return 0, err + } + if err := writeAll(r.f, opt); err != nil { + return 0, err + } + if err := r.writeTrailerLocked(length); err != nil { + return 0, err + } + id := r.ifaces + r.ifaces++ + return id, nil +} + +func (r *Run) NewSession(network string, local, remote netip.AddrPort, iface int) (*Session, error) { + if r == nil { + return nil, nil + } + return &Session{run: r, netw: network, local: local, remote: remote, iface: iface}, nil +} + +func (r *Run) Close() error { + if r == nil { + return nil + } + r.mu.Lock() + defer r.mu.Unlock() + if r.f == nil { + return nil + } + err := r.f.Close() + r.f = nil + return err +} + +func (r *Run) Stats() (packets uint64, first, last time.Time) { + if r == nil { + return 0, time.Time{}, time.Time{} + } + r.mu.Lock() + defer r.mu.Unlock() + return r.packets, r.first, r.last +} + +func (r *Run) writeSHB() error { + opt := encodeOption(2, []byte("GoNetSim simulated network")) + opt = append(opt, encodeOption(3, []byte(runtime.GOOS+"/"+runtime.GOARCH))...) + opt = append(opt, encodeOption(4, []byte("gonetsim"))...) + opt = append(opt, encodeOption(0, nil)...) + length := 28 + len(opt) + b := make([]byte, 24) + binary.LittleEndian.PutUint32(b[0:4], 0x0A0D0D0A) + binary.LittleEndian.PutUint32(b[4:8], uint32(length)) + binary.LittleEndian.PutUint32(b[8:12], 0x1A2B3C4D) + binary.LittleEndian.PutUint16(b[12:14], 1) + binary.LittleEndian.PutUint16(b[14:16], 0) + binary.LittleEndian.PutUint64(b[16:24], 0xFFFFFFFFFFFFFFFF) + if _, err := r.f.Write(b); err != nil { + return err + } + if _, err := r.f.Write(opt); err != nil { + return err + } + return r.writeTrailerLocked(length) +} + +func (r *Run) writeTrailerLocked(length int) error { + var b [4]byte + binary.LittleEndian.PutUint32(b[:], uint32(length)) + _, err := r.f.Write(b[:]) + return err +} diff --git a/internal/capture/session.go b/internal/capture/session.go new file mode 100644 index 0000000..6de4452 --- /dev/null +++ b/internal/capture/session.go @@ -0,0 +1,303 @@ +package capture + +import ( + "encoding/binary" + "net" + "net/netip" + "os" + "time" + + "github.com/google/gopacket" + "github.com/google/gopacket/layers" +) + +const ( + snapLen = 262144 +) + +type Session struct { + run *Run + iface int + netw string // "tcp" or "udp" + local netip.AddrPort + remote netip.AddrPort + + synSent bool + pending string + + clientSeq uint32 + serverSeq uint32 +} + +func (s *Session) Comment(text string) { + if s == nil || s.run == nil { + return + } + s.run.mu.Lock() + defer s.run.mu.Unlock() + s.pending = text +} + +func (s *Session) Write(data []byte, fromClient bool) error { + if s == nil || s.run == nil { + return nil + } + s.run.mu.Lock() + defer s.run.mu.Unlock() + if s.netw == "udp" { + return s.writeUDP(data, fromClient) + } + return s.writeTCP(data, fromClient) +} + +func (s *Session) Close() error { + if s == nil || s.run == nil { + return nil + } + s.run.mu.Lock() + defer s.run.mu.Unlock() + if s.netw == "tcp" && s.synSent { + if err := s.emitTCP(true, false, true, true, s.clientSeq, s.serverSeq); err != nil { + return err + } + s.clientSeq++ + if err := s.emitTCP(false, false, true, true, s.serverSeq, s.clientSeq); err != nil { + return err + } + } + s.synSent = false + return nil +} + +func (s *Session) Flush() error { + if s == nil || s.run == nil { + return nil + } + s.run.mu.Lock() + defer s.run.mu.Unlock() + return s.run.f.Sync() +} + +func encodeOption(code uint16, value []byte) []byte { + out := make([]byte, 4+len(value)) + binary.LittleEndian.PutUint16(out[0:2], code) + binary.LittleEndian.PutUint16(out[2:4], uint16(len(value))) + copy(out[4:], value) + for len(out)%4 != 0 { + out = append(out, 0) + } + return out +} + +func (s *Session) writeUDP(data []byte, fromClient bool) error { + src, dst := s.endpoints(fromClient) + _, err := s.epb(s.build(data, src, dst, isUDP)) + return err +} + +func (s *Session) writeTCP(data []byte, fromClient bool) error { + if !s.synSent { + if _, err := s.epb(s.buildTCPControl(true, true, false, false, s.clientSeq, s.serverSeq)); err != nil { + return err + } + s.synSent = true + s.clientSeq++ + } + if s.serverSeq == 0 { + if _, err := s.epb(s.buildTCPControl(false, true, true, false, s.serverSeq, s.clientSeq)); err != nil { + return err + } + s.serverSeq++ + } + + seq, ack := s.clientSeq, s.serverSeq + if !fromClient { + seq, ack = s.serverSeq, s.clientSeq + } + _, err := s.epb(s.buildTCPData(data, fromClient, seq, ack)) + if fromClient { + s.clientSeq += uint32(len(data)) + } else { + s.serverSeq += uint32(len(data)) + } + return err + +} + +func (s *Session) emitTCP(fromClient, syn, ackFlag, fin bool, seq, ackNum uint32) error { + _, err := s.epb(s.buildTCPControl(fromClient, syn, ackFlag, fin, seq, ackNum)) + return err +} + +func (s *Session) epb(frame []byte) (int, error) { + ts := time.Now() + opts := s.takeComment() + length := 32 + frameLen(frame) + len(opts) + b := make([]byte, 28) + binary.LittleEndian.PutUint32(b[0:4], 6) + binary.LittleEndian.PutUint32(b[4:8], uint32(length)) + binary.LittleEndian.PutUint32(b[8:12], uint32(s.iface)) + binary.LittleEndian.PutUint32(b[12:16], uint32(ts.UnixMicro()>>32)) + binary.LittleEndian.PutUint32(b[16:20], uint32(ts.UnixMicro())) + binary.LittleEndian.PutUint32(b[20:24], uint32(len(frame))) + binary.LittleEndian.PutUint32(b[24:28], uint32(len(frame))) + if err := writeAll(s.run.f, b); err != nil { + return 0, err + } + if err := writeAll(s.run.f, frame); err != nil { + return 0, err + } + if pad := framePad(len(frame)); pad > 0 { + if err := writeAll(s.run.f, make([]byte, pad)); err != nil { + return 0, err + } + } + if err := writeAll(s.run.f, opts); err != nil { + return 0, err + } + if err := s.run.writeTrailerLocked(length); err != nil { + return 0, err + } + s.run.packets++ + if s.run.packets == 1 { + s.run.first = ts + } + s.run.last = ts + return len(frame), nil +} + +func (s *Session) takeComment() []byte { + if s.pending == "" { + return nil + } + out := encodeOption(1, []byte(s.pending)) // opt_comment, no nul + s.pending = "" + return out +} + +func frameLen(b []byte) int { + return len(b) + framePad(len(b)) +} + +func framePad(n int) int { + return (4 - n%4) % 4 +} + +func writeAll(f *os.File, b []byte) error { + _, err := f.Write(b) + return err +} + +type transportKind int + +const ( + isUDP transportKind = iota + isTCP +) + +func (s *Session) build(data []byte, src, dst netip.AddrPort, kind transportKind) []byte { + var network, transport gopacket.SerializableLayer + + switch kind { + case isTCP: + tcp := &layers.TCP{SrcPort: layers.TCPPort(src.Port()), DstPort: layers.TCPPort(dst.Port())} + network = tcpIPLayer(src, dst, layers.IPProtocolTCP, tcp) + transport = tcp + default: + udp := &layers.UDP{SrcPort: layers.UDPPort(src.Port()), DstPort: layers.UDPPort(dst.Port())} + network = udpIPLayer(src, dst, udp) + transport = udp + } + return s.serialize(data, src, network, transport) +} + +func (s *Session) buildTCPData(data []byte, fromClient bool, seq, ack uint32) []byte { + src, dst := s.endpoints(fromClient) + tcp := newTCPLayer(src, dst, seq, ack, false, true, false, len(data) > 0) + return s.buildWithTCP(data, src, dst, tcp) +} + +func (s *Session) buildTCPControl(fromClient, syn, ackFlag, fin bool, seq, ackNum uint32) []byte { + src, dst := s.endpoints(fromClient) + tcp := newTCPLayer(src, dst, seq, ackNum, syn, ackFlag, fin, false) + return s.buildWithTCP(nil, src, dst, tcp) +} + +func newTCPLayer(src, dst netip.AddrPort, seq, ack uint32, syn, ackFlag, fin, psh bool) *layers.TCP { + return &layers.TCP{ + SrcPort: layers.TCPPort(src.Port()), + DstPort: layers.TCPPort(dst.Port()), + Seq: seq, + Ack: ack, + SYN: syn, + ACK: ackFlag, + FIN: fin, + PSH: psh, + Window: 65535, + } +} + +func (s *Session) buildWithTCP(data []byte, src, dst netip.AddrPort, tcp *layers.TCP) []byte { + return s.serialize(data, src, tcpIPLayer(src, dst, layers.IPProtocolTCP, tcp), tcp) +} + +func tcpIPLayer(src, dst netip.AddrPort, proto layers.IPProtocol, tcp *layers.TCP) gopacket.SerializableLayer { + if src.Addr().Is4() { + ip := &layers.IPv4{Version: 4, TTL: 64, Protocol: proto, SrcIP: src.Addr().AsSlice(), DstIP: dst.Addr().AsSlice()} + _ = tcp.SetNetworkLayerForChecksum(ip) + return ip + } + ip := &layers.IPv6{Version: 6, HopLimit: 64, NextHeader: proto, SrcIP: src.Addr().AsSlice(), DstIP: dst.Addr().AsSlice()} + _ = tcp.SetNetworkLayerForChecksum(ip) + return ip +} + +func udpIPLayer(src, dst netip.AddrPort, udp *layers.UDP) gopacket.SerializableLayer { + if src.Addr().Is4() { + ip := &layers.IPv4{Version: 4, TTL: 64, Protocol: layers.IPProtocolUDP, SrcIP: src.Addr().AsSlice(), DstIP: dst.Addr().AsSlice()} + _ = udp.SetNetworkLayerForChecksum(ip) + return ip + } + ip := &layers.IPv6{Version: 6, HopLimit: 64, NextHeader: layers.IPProtocolUDP, SrcIP: src.Addr().AsSlice(), DstIP: dst.Addr().AsSlice()} + _ = udp.SetNetworkLayerForChecksum(ip) + return ip +} + +func (s *Session) serialize(data []byte, src netip.AddrPort, network, transport gopacket.SerializableLayer) []byte { + eth := s.ethLayer(src) + buf := gopacket.NewSerializeBuffer() + layersToWrite := []gopacket.SerializableLayer{ð, network, transport} + if len(data) > 0 { + layersToWrite = append(layersToWrite, gopacket.Payload(data)) + } + if err := gopacket.SerializeLayers(buf, serializeOpts, layersToWrite...); err != nil { + return nil + } + return buf.Bytes() +} + +func (s *Session) ethLayer(src netip.AddrPort) layers.Ethernet { + var etherType = layers.EthernetTypeIPv4 + if !src.Addr().Is4() { + etherType = layers.EthernetTypeIPv6 + } + eth := layers.Ethernet{SrcMAC: clientMAC, DstMAC: serverMAC, EthernetType: etherType} + if src.Addr() == s.local.Addr() { + eth.SrcMAC = serverMAC + eth.DstMAC = clientMAC + } + return eth +} + +func (s *Session) endpoints(fromClient bool) (netip.AddrPort, netip.AddrPort) { + if fromClient { + return s.remote, s.local + } + return s.local, s.remote +} + +var ( + serializeOpts = gopacket.SerializeOptions{FixLengths: true, ComputeChecksums: true} + clientMAC = net.HardwareAddr{0x02, 0x00, 0x00, 0x00, 0x00, 0x01} + serverMAC = net.HardwareAddr{0x02, 0x00, 0x00, 0x00, 0x00, 0x02} +) diff --git a/internal/capture/session_test.go b/internal/capture/session_test.go new file mode 100644 index 0000000..2a80799 --- /dev/null +++ b/internal/capture/session_test.go @@ -0,0 +1,451 @@ +////---------------------------------------------------------------------------- +// NOTICE: to save development time, test files (including this) have been +// generated with LLMs. The author(s) do not claim credit for these tests +// and exist purely for maximising code quality and reliability +// +// For more information please see `/.github/AI_USAGE.md` +//----------------------------------------------------------------------------// + +package capture + +import ( + "bytes" + "encoding/binary" + "io" + "net/netip" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/google/gopacket" + "github.com/google/gopacket/layers" + "github.com/google/gopacket/pcapgo" +) + +type frame struct { + src, dst netip.AddrPort + syn, ack, fin bool + seq, ackNum uint32 + payload string +} + +func testRun(t *testing.T) (*Run, string) { + t.Helper() + path := filepath.Join(t.TempDir(), "run.pcapng") + run, err := NewRun(path) + if err != nil { + t.Fatalf("NewRun: %v", err) + } + t.Cleanup(func() { _ = run.Close() }) + if _, err := run.NewInterface("test"); err != nil { + t.Fatalf("NewInterface: %v", err) + } + return run, path +} + +func testSession(t *testing.T, run *Run, network string, local, remote netip.AddrPort) *Session { + t.Helper() + ses, err := run.NewSession(network, local, remote, 0) + if err != nil { + t.Fatalf("NewSession: %v", err) + } + return ses +} + +func readFrames(t *testing.T, path string) []gopacket.Packet { + t.Helper() + f, err := os.Open(path) + if err != nil { + t.Fatalf("Open: %v", err) + } + defer f.Close() + r, err := pcapgo.NewNgReader(f, pcapgo.DefaultNgReaderOptions) + if err != nil { + t.Fatalf("NewNgReader: %v", err) + } + if r.LinkType() != layers.LinkTypeEthernet { + t.Fatalf("LinkType = %v, want Ethernet", r.LinkType()) + } + var out []gopacket.Packet + for { + data, _, err := r.ReadPacketData() + if err == io.EOF { + break + } + if err != nil { + t.Fatalf("ReadPacketData: %v", err) + } + p := gopacket.NewPacket(data, layers.LinkTypeEthernet, gopacket.Default) + if p.ErrorLayer() != nil { + t.Fatalf("packet failed to decode: %v", p.ErrorLayer().Error()) + } + out = append(out, p) + } + return out +} + +func checkTCP(t *testing.T, p gopacket.Packet, want frame) { + t.Helper() + tcp, ok := p.Layer(layers.LayerTypeTCP).(*layers.TCP) + if !ok { + t.Fatalf("packet is not TCP: %v", p.Layers()) + } + if tcp.SrcPort != layers.TCPPort(want.src.Port()) || tcp.DstPort != layers.TCPPort(want.dst.Port()) { + t.Errorf("ports = %s:%s, want %d:%d", tcp.SrcPort, tcp.DstPort, want.src.Port(), want.dst.Port()) + } + if tcp.SYN != want.syn || tcp.ACK != want.ack || tcp.FIN != want.fin { + t.Errorf("flags SYN=%v ACK=%v FIN=%v, want SYN=%v ACK=%v FIN=%v", tcp.SYN, tcp.ACK, tcp.FIN, want.syn, want.ack, want.fin) + } + if tcp.Seq != want.seq || tcp.Ack != want.ackNum { + t.Errorf("seq/ack = %d/%d, want %d/%d", tcp.Seq, tcp.Ack, want.seq, want.ackNum) + } + if string(tcp.Payload) != want.payload { + t.Errorf("payload = %q, want %q", tcp.Payload, want.payload) + } +} + +func TestSessionTCP(t *testing.T) { + local := netip.MustParseAddrPort("127.0.0.1:8080") + remote := netip.MustParseAddrPort("10.0.0.5:40000") + run, path := testRun(t) + + ses := testSession(t, run, "tcp", local, remote) + if err := ses.Write([]byte("hello"), true); err != nil { + t.Fatalf("Write client: %v", err) + } + if err := ses.Write([]byte("world"), false); err != nil { + t.Fatalf("Write server: %v", err) + } + if err := ses.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + pkts := readFrames(t, path) + want := []frame{ + {remote, local, true, false, false, 0, 0, ""}, + {local, remote, true, true, false, 0, 1, ""}, + {remote, local, false, true, false, 1, 1, "hello"}, + {local, remote, false, true, false, 1, 6, "world"}, + {remote, local, false, true, true, 6, 6, ""}, + {local, remote, false, true, true, 6, 7, ""}, + } + if len(pkts) != len(want) { + t.Fatalf("got %d packets, want %d (SYN, SYN-ACK, 2 data, 2 FIN)", len(pkts), len(want)) + } + for i, w := range want { + checkTCP(t, pkts[i], w) + } +} + +func TestSessionTCPIPv6(t *testing.T) { + local := netip.MustParseAddrPort("[::1]:8080") + remote := netip.MustParseAddrPort("[2001:db8::5]:40000") + run, path := testRun(t) + + ses := testSession(t, run, "tcp", local, remote) + if err := ses.Write([]byte("ping"), true); err != nil { + t.Fatalf("Write: %v", err) + } + if err := ses.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + pkts := readFrames(t, path) + if len(pkts) != 5 { + t.Fatalf("got %d packets, want 5", len(pkts)) + } + if pkts[0].Layer(layers.LayerTypeIPv6) == nil { + t.Fatalf("expected IPv6 frames, got %v", pkts[0].Layers()) + } +} + +func TestSessionUDP(t *testing.T) { + local := netip.MustParseAddrPort("127.0.0.1:12345") + remote := netip.MustParseAddrPort("10.0.0.5:5000") + run, path := testRun(t) + + ses := testSession(t, run, "udp", local, remote) + if err := ses.Write([]byte("query"), true); err != nil { + t.Fatalf("Write client: %v", err) + } + if err := ses.Write([]byte("answer"), false); err != nil { + t.Fatalf("Write server: %v", err) + } + if err := ses.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + pkts := readFrames(t, path) + if len(pkts) != 2 { + t.Fatalf("got %d packets, want 2", len(pkts)) + } + for i, want := range []struct { + src, dst netip.AddrPort + payload string + }{ + {remote, local, "query"}, + {local, remote, "answer"}, + } { + udp, ok := pkts[i].Layer(layers.LayerTypeUDP).(*layers.UDP) + if !ok { + t.Fatalf("packet %d is not UDP", i) + } + if udp.SrcPort != layers.UDPPort(want.src.Port()) || udp.DstPort != layers.UDPPort(want.dst.Port()) { + t.Errorf("packet %d ports = %s:%s, want %d:%d", i, udp.SrcPort, udp.DstPort, want.src.Port(), want.dst.Port()) + } + if string(udp.Payload) != want.payload { + t.Errorf("packet %d payload = %q, want %q", i, udp.Payload, want.payload) + } + } +} + +func TestSessionComment(t *testing.T) { + local := netip.MustParseAddrPort("127.0.0.1:8080") + remote := netip.MustParseAddrPort("10.0.0.5:40000") + run, path := testRun(t) + + ses := testSession(t, run, "tcp", local, remote) + ses.Comment("he-lo") + if err := ses.Write([]byte("hello"), true); err != nil { + t.Fatalf("Write: %v", err) + } + if err := ses.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + raw, err := os.ReadFile(path) + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + // the comment attaches to the next frame (here SYN); the file must still + // parse cleanly as pcapng + if !bytes.Contains(raw, []byte("he-lo")) { + t.Fatal("comment text not found in pcapng bytes") + } + if pkts := readFrames(t, path); len(pkts) != 5 { + t.Fatalf("got %d packets, want 5", len(pkts)) + } +} + +func TestSessionEmpty(t *testing.T) { + local := netip.MustParseAddrPort("127.0.0.1:8080") + remote := netip.MustParseAddrPort("10.0.0.5:40000") + run, path := testRun(t) + + ses := testSession(t, run, "tcp", local, remote) + if err := ses.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if pkts := readFrames(t, path); len(pkts) != 0 { + t.Fatalf("expected empty capture, got %d packets", len(pkts)) + } +} + +func TestRun(t *testing.T) { + t.Run("one file holds many flows", func(t *testing.T) { + run, path := testRun(t) + local := netip.MustParseAddrPort("127.0.0.1:53") + remote := netip.MustParseAddrPort("203.0.113.10:43210") + + udp := testSession(t, run, "udp", local, remote) + if err := udp.Write([]byte("query"), true); err != nil { + t.Fatalf("Write: %v", err) + } + if err := udp.Write([]byte("answer"), false); err != nil { + t.Fatalf("Write: %v", err) + } + if err := udp.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + tcp := testSession(t, run, "tcp", local, remote) + if err := tcp.Write([]byte("hello"), true); err != nil { + t.Fatalf("Write: %v", err) + } + if err := tcp.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + info, err := Inspect(path) + if err != nil { + t.Fatalf("Inspect: %v", err) + } + if info.LinkType != layers.LinkTypeEthernet || info.Packets != 7 { + t.Fatalf("unexpected inspect result %+v", info) + } + if len(info.Interfaces) != 1 || info.Interfaces[0] != "test" { + t.Fatalf("unexpected interfaces %+v", info.Interfaces) + } + if info.CreatedBy != "gonetsim" { + t.Fatalf("unexpected created-by %q", info.CreatedBy) + } + if packets, _, _ := run.Stats(); packets != 7 { + t.Fatalf("Stats packets = %d, want 7", packets) + } + }) + + t.Run("run path resolution", func(t *testing.T) { + dir := t.TempDir() + t.Setenv("XDG_DATA_HOME", dir) + got, err := RunPath("") + if err != nil { + t.Fatalf("RunPath: %v", err) + } + wantDir := filepath.Join(dir, "gonetsim", "runs") + if filepath.Dir(got) != wantDir || !strings.HasSuffix(got, ".pcapng") { + t.Fatalf("RunPath = %q, want dir %q with .pcapng suffix", got, wantDir) + } + + explicit := filepath.Join(dir, "case", "run.pcapng") + got, err = RunPath(explicit) + if err != nil { + t.Fatalf("RunPath explicit: %v", err) + } + if got != explicit { + t.Fatalf("RunPath explicit = %q, want %q", got, explicit) + } + if st, err := os.Stat(filepath.Join(dir, "case")); err != nil || !st.IsDir() { + t.Fatalf("expected parent dir to be created: %v", err) + } + }) + + t.Run("run id format", func(t *testing.T) { + id := NewRunID() + if len(id) != 20 || id[8] != '-' || id[15] != '-' { + t.Fatalf("unexpected run id %q", id) + } + }) + + t.Run("empty run inspects cleanly", func(t *testing.T) { + run, path := testRun(t) + if err := run.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + info, err := Inspect(path) + if err != nil { + t.Fatalf("Inspect: %v", err) + } + if info.Packets != 0 { + t.Fatalf("Packets = %d, want 0", info.Packets) + } + }) + + t.Run("nil run is a no-op", func(t *testing.T) { + var run *Run + if run.Path() != "" { + t.Fatalf("expected empty path from nil run") + } + if packets, first, last := run.Stats(); packets != 0 || !first.IsZero() || !last.IsZero() { + t.Fatalf("expected zero stats from nil run") + } + if err := run.Close(); err != nil { + t.Fatalf("Close on nil run: %v", err) + } + if iface, err := run.NewInterface("x"); err != nil || iface != 0 { + t.Fatalf("expected zero interface from nil run, got %d, %v", iface, err) + } + ses, err := run.NewSession("tcp", + netip.MustParseAddrPort("5.6.7.8:9"), netip.MustParseAddrPort("1.2.3.4:5"), 0) + if err != nil || ses != nil { + t.Fatalf("expected nil session from nil run, got %v, %v", ses, err) + } + ses.Comment("ignored") + _ = ses.Write([]byte("data"), true) + if err := ses.Close(); err != nil { + t.Fatalf("Close on nil session: %v", err) + } + }) +} + +func TestInspect(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "inspect", "manual.pcapng") + if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { + t.Fatalf("MkdirAll: %v", err) + } + f, err := os.Create(path) + if err != nil { + t.Fatalf("Create: %v", err) + } + w, err := pcapgo.NewNgWriter(f, layers.LinkTypeEthernet) + if err != nil { + t.Fatalf("NewNgWriter: %v", err) + } + ts := time.Date(2026, 1, 2, 3, 4, 5, 0, time.UTC) + for i, p := range [][]byte{{0xde, 0xad}, {0xca, 0xfe}} { + ci := gopacket.CaptureInfo{Timestamp: ts.Add(time.Duration(i) * time.Second), CaptureLength: len(p), Length: len(p)} + if err := w.WritePacket(ci, p); err != nil { + t.Fatalf("WritePacket: %v", err) + } + } + if err := w.Flush(); err != nil { + t.Fatalf("Flush: %v", err) + } + if err := f.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + info, err := Inspect(path) + if err != nil { + t.Fatalf("Inspect: %v", err) + } + if info.LinkType != layers.LinkTypeEthernet { + t.Fatalf("LinkType = %v, want Ethernet", info.LinkType) + } + if info.Packets != 2 { + t.Fatalf("Packets = %d, want 2", info.Packets) + } + if !info.First.Equal(ts) || !info.Last.Equal(ts.Add(time.Second)) { + t.Fatalf("First/Last timestamps = %v/%v, want %v/%v", info.First, info.Last, ts, ts.Add(time.Second)) + } + + legacy := filepath.Join(dir, "legacy.pcap") + if err := os.WriteFile(legacy, []byte{0xd4, 0xc3, 0xb2, 0xa1}, 0o644); err != nil { + t.Fatalf("WriteFile: %v", err) + } + if _, err := Inspect(legacy); err == nil || !strings.Contains(err.Error(), "legacy pcap") { + t.Fatalf("expected legacy pcap error, got %v", err) + } +} + +func TestBlockLengthsMatch(t *testing.T) { + local := netip.MustParseAddrPort("127.0.0.1:8080") + remote := netip.MustParseAddrPort("10.0.0.5:40000") + run, path := testRun(t) + + ses := testSession(t, run, "tcp", local, remote) + ses.Comment("greeting") + if err := ses.Write([]byte("hello"), true); err != nil { + t.Fatalf("Write: %v", err) + } + if err := ses.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + raw, err := os.ReadFile(path) + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + off := 0 + blocks := 0 + for off < len(raw) { + if len(raw)-off < 12 { + t.Fatalf("block %d at offset %d: truncated header", blocks, off) + } + blen := int(binary.LittleEndian.Uint32(raw[off+4 : off+8])) + if blen < 12 || off+blen > len(raw) { + t.Fatalf("block %d at offset %d: bad length %d (file %d)", blocks, off, blen, len(raw)) + } + trail := binary.LittleEndian.Uint32(raw[off+blen-4 : off+blen]) + if int(trail) != blen { + t.Fatalf("block %d at offset %d: lengths %d and %d don't match", blocks, off, blen, trail) + } + off += blen + blocks++ + } + if blocks == 0 { + t.Fatalf("no blocks found") + } +} diff --git a/internal/config/config.go b/internal/config/config.go index bed63bb..bb3d80a 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -14,8 +14,6 @@ import ( "github.com/knadh/koanf/providers/confmap" "github.com/knadh/koanf/providers/file" "github.com/knadh/koanf/v2" - - "github.com/lachlanharrisdev/gonetsim/internal/state" ) const ( @@ -45,34 +43,37 @@ type GeneralConfig struct { ShutdownTimeout time.Duration `koanf:"shutdown_timeout"` } +type ServiceBase struct { + Enabled bool `koanf:"enabled"` + Listen string `koanf:"listen"` + Capture bool `koanf:"capture"` +} + type DNSConfig struct { - Enabled bool `koanf:"enabled"` - Listen string `koanf:"listen"` - Network string `koanf:"network"` - IPv4 string `koanf:"ipv4"` - IPv6 string `koanf:"ipv6"` - Domain string `koanf:"domain"` - TXT string `koanf:"txt"` - TTL uint32 `koanf:"ttl"` - Compress bool `koanf:"compress"` + ServiceBase `koanf:",squash"` + Network string `koanf:"network"` + IPv4 string `koanf:"ipv4"` + IPv6 string `koanf:"ipv6"` + Domain string `koanf:"domain"` + TXT string `koanf:"txt"` + TTL uint32 `koanf:"ttl"` + Compress bool `koanf:"compress"` } type HTTPConfig struct { - Enabled bool `koanf:"enabled"` - Listen string `koanf:"listen"` - Status int `koanf:"status"` - Mode string `koanf:"mode"` - RootDir string `koanf:"root_dir"` + ServiceBase `koanf:",squash"` + Status int `koanf:"status"` + Mode string `koanf:"mode"` + RootDir string `koanf:"root_dir"` } type HTTPSConfig struct { - Enabled bool `koanf:"enabled"` - Listen string `koanf:"listen"` - Status int `koanf:"status"` - Mode string `koanf:"mode"` - RootDir string `koanf:"root_dir"` - Cert string `koanf:"cert"` - Key string `koanf:"key"` + ServiceBase `koanf:",squash"` + Status int `koanf:"status"` + Mode string `koanf:"mode"` + RootDir string `koanf:"root_dir"` + Cert string `koanf:"cert"` + Key string `koanf:"key"` } type LoggingConfig struct { @@ -107,27 +108,24 @@ func Default() Config { return Config{ General: GeneralConfig{ShutdownTimeout: 2 * time.Second}, DNS: DNSConfig{ - Enabled: true, - Listen: ":53", - Network: "udp", - IPv4: "auto", - IPv6: "::1", - Domain: "localhost", - TXT: "TXT record response from GoNetSim", - TTL: 60, - Compress: false, + ServiceBase: ServiceBase{Enabled: true, Listen: ":53", Capture: true}, + Network: "udp", + IPv4: "auto", + IPv6: "::1", + Domain: "localhost", + TXT: "TXT record response from GoNetSim", + TTL: 60, + Compress: false, }, HTTP: HTTPConfig{ - Enabled: true, - Listen: ":80", - Status: 200, - Mode: "fake", + ServiceBase: ServiceBase{Enabled: true, Listen: ":80", Capture: true}, + Status: 200, + Mode: "fake", }, HTTPS: HTTPSConfig{ - Enabled: true, - Listen: ":443", - Status: 200, - Mode: "fake", + ServiceBase: ServiceBase{Enabled: true, Listen: ":443", Capture: true}, + Status: 200, + Mode: "fake", }, Logging: LoggingConfig{ LogFormat: "text", @@ -162,7 +160,7 @@ func (c Config) Validate() error { } if strings.TrimSpace(c.State.TotalLimit) != "" { - if _, err := state.ParseSize(c.State.TotalLimit); err != nil { + if _, err := ParseSize(c.State.TotalLimit); err != nil { return fmt.Errorf("state.total_limit: %w", err) } } @@ -192,12 +190,6 @@ func LoadOrCreate(configPath string) (LoadResult, error) { return LoadOrCreateWithOverrides(configPath, nil) } -// LoadOrCreateWithOverrides loads defaults, then the on-disk config file, then applies -// the provided flat overrides (dot-delimited keys). -// -// Validation is intentionally not run here; callers should map the resulting config -// into the isolated service configs (e.g. dnsserver.Config) and call Validate() once -// on those structs before starting services. func LoadOrCreateWithOverrides(configPath string, overrides map[string]any) (LoadResult, error) { resolved, created, err := resolveAndCreate(configPath) if err != nil { diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 650b8ca..cf9aca8 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -1,3 +1,11 @@ +////---------------------------------------------------------------------------- +// NOTICE: to save development time, test files (including this) have been +// generated with LLMs. The author(s) do not claim credit for these tests +// and exist purely for maximising code quality and reliability +// +// For more information please see `/.github/AI_USAGE.md` +//----------------------------------------------------------------------------// + package config import ( @@ -7,9 +15,6 @@ import ( "time" ) -// / -// / verifies that a new config file is created when one doesn't exist, checks loading, & that a second call doesn't overwrite the file -// / func TestLoadOrCreate_CreatesAndLoadsConfig(t *testing.T) { dir := t.TempDir() path := filepath.Join(dir, "gonetsim.toml") @@ -183,3 +188,29 @@ func TestFirstExistingFile_PrefersLocalThenUserThenSystem(t *testing.T) { t.Fatalf("expected first (highest precedence) file %q, got %q", a, got) } } + +func TestLegacyCaptureDirIgnored(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "gonetsim.toml") + content := ` +[http] +enabled = true +listen = "127.0.0.1:0" +capture = true +capture_dir = "/tmp/should-be-ignored" +` + if err := os.WriteFile(path, []byte(content), 0o644); err != nil { + t.Fatalf("WriteFile: %v", err) + } + + res, err := LoadOrCreate(path) + if err != nil { + t.Fatalf("LoadOrCreate: %v", err) + } + if err := res.Config.Validate(); err != nil { + t.Fatalf("Validate: %v", err) + } + if !res.Config.HTTP.Capture { + t.Fatalf("expected http.capture to survive") + } +} diff --git a/internal/config/default_config.toml b/internal/config/default_config.toml index c390701..8b8d0ea 100644 --- a/internal/config/default_config.toml +++ b/internal/config/default_config.toml @@ -36,6 +36,9 @@ txt = "TXT record response from GoNetSim" ttl = 60 # Enable DNS message compression compress = false +# Save every query/response flow to the run capture file (udp and/or tcp, +# matching the configured network). +capture = true [http] enabled = true @@ -46,6 +49,9 @@ status = 200 # directory specified in root_dir (required when mode = "real") mode = "fake" root_dir = "" +# Save every connection to the run capture file. +# Inspect it with `gonetsim pcap `. +capture = true [https] enabled = true @@ -56,6 +62,9 @@ status = 200 # directory specified in root_dir (required when mode = "real") mode = "fake" root_dir = "" +# Save every connection to the run capture file. +# Note: captures hold TLS ciphertext, not plaintext. +capture = true [logging] # Log output format: "text" or "json". @@ -94,7 +103,7 @@ total_limit = "64MiB" # tls = false # tls_cert = "" # tls_key = "" -# # Write everything a client sends to the artifacts directory. +# # Write everything a client sends to the run capture file. # capture = true # SMTP-style mail sink, served by the example Lua handler: diff --git a/internal/config/size.go b/internal/config/size.go new file mode 100644 index 0000000..aaa7239 --- /dev/null +++ b/internal/config/size.go @@ -0,0 +1,41 @@ +package config + +import ( + "fmt" + "strconv" + "strings" +) + +func ParseSize(s string) (int64, error) { + s = strings.TrimSpace(strings.ToLower(s)) + if n, err := strconv.ParseInt(s, 10, 64); err == nil { + if n <= 0 { + return 0, fmt.Errorf("size must be positive") + } + return n, nil + } + + var mult int64 + switch { + case strings.HasSuffix(s, "kib"): + mult, s = 1<<10, s[:len(s)-3] + case strings.HasSuffix(s, "mib"): + mult, s = 1<<20, s[:len(s)-3] + case strings.HasSuffix(s, "gib"): + mult, s = 1<<30, s[:len(s)-3] + case strings.HasSuffix(s, "k"): + mult, s = 1<<10, s[:len(s)-1] + case strings.HasSuffix(s, "m"): + mult, s = 1<<20, s[:len(s)-1] + case strings.HasSuffix(s, "g"): + mult, s = 1<<30, s[:len(s)-1] + default: + return 0, fmt.Errorf("invalid size %q (expected e.g. 64MiB)", s) + } + + n, err := strconv.ParseInt(strings.TrimSpace(s), 10, 64) + if err != nil || n <= 0 { + return 0, fmt.Errorf("invalid size %q (expected e.g. 64MiB)", s) + } + return n * mult, nil +} diff --git a/internal/config/size_test.go b/internal/config/size_test.go new file mode 100644 index 0000000..2a9ea8b --- /dev/null +++ b/internal/config/size_test.go @@ -0,0 +1,31 @@ +package config + +import "testing" + +func TestParseSize(t *testing.T) { + cases := []struct { + in string + want int64 + wantErr bool + }{ + {"64MiB", 64 << 20, false}, + {"64mib", 64 << 20, false}, + {"512K", 512 << 10, false}, + {"1GiB", 1 << 30, false}, + {"4096", 4096, false}, + {"", 0, true}, + {"64GiB", 64 << 30, false}, + {"abc", 0, true}, + {"-1MiB", 0, true}, + {"64TiB", 0, true}, + } + for _, tc := range cases { + got, err := ParseSize(tc.in) + if tc.wantErr && err == nil { + t.Errorf("ParseSize(%q): expected error", tc.in) + } + if !tc.wantErr && (err != nil || got != tc.want) { + t.Errorf("ParseSize(%q) = %d, %v; want %d", tc.in, got, err, tc.want) + } + } +} diff --git a/internal/dnsserver/capture_test.go b/internal/dnsserver/capture_test.go new file mode 100644 index 0000000..c6c0d1b --- /dev/null +++ b/internal/dnsserver/capture_test.go @@ -0,0 +1,215 @@ +// //---------------------------------------------------------------------------- +// // NOTICE: to save development time, test files (including this) have been +// // generated with LLMs. The author(s) do not claim credit for these tests +// // and exist purely for maximising code quality and reliability +// // +// // For more information please see `/.github/AI_USAGE.md` +// //----------------------------------------------------------------------------// + +package dnsserver + +import ( + "context" + "fmt" + "net" + "net/netip" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/google/gopacket" + "github.com/google/gopacket/layers" + "github.com/google/gopacket/pcapgo" + "github.com/miekg/dns" + + "github.com/lachlanharrisdev/gonetsim/internal/capture" + "github.com/lachlanharrisdev/gonetsim/internal/service" + "github.com/lachlanharrisdev/gonetsim/internal/testutil" +) + +func testRun(t *testing.T) (*capture.Run, string) { + t.Helper() + path := filepath.Join(t.TempDir(), "run.pcapng") + run, err := capture.NewRun(path) + if err != nil { + t.Fatalf("NewRun: %v", err) + } + t.Cleanup(func() { _ = run.Close() }) + return run, path +} + +func TestService_CapturesUDP(t *testing.T) { + conf := baseCaptureConfig(t, "udp") + conf.Capture = true + run, path := testRun(t) + + svc, errCh := startDNSService(t, conf, run) + + query := newAQuery() + client := &dns.Client{Net: "udp", Timeout: 1 * time.Second} + _, _, err := retryExchange(t, client, conf.Addr, query) + if err != nil { + t.Fatalf("exchange: %v", err) + } + + waitDNSCapture(t, path, "example") + waitTransportPayloads(t, path, func(s string) bool { + return strings.Count(s, "example") >= 2 // query + response + }) + + svc.Stop(context.Background()) //nolint:errcheck,gosec + discardStartErr(t, errCh) +} + +func TestService_CapturesTCP(t *testing.T) { + conf := baseCaptureConfig(t, "tcp") + conf.Capture = true + run, path := testRun(t) + + svc, errCh := startDNSService(t, conf, run) + + query := newAQuery() + client := &dns.Client{Net: "tcp", Timeout: 1 * time.Second} + _, _, err := retryExchange(t, client, conf.Addr, query) + if err != nil { + t.Fatalf("exchange: %v", err) + } + + waitDNSCapture(t, path, "example") + waitTransportPayloads(t, path, func(s string) bool { + return strings.Count(s, "example") >= 2 // query + response + }) + + svc.Stop(context.Background()) //nolint:errcheck,gosec + discardStartErr(t, errCh) +} + +func baseCaptureConfig(t *testing.T, network string) Config { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("Listen: %v", err) + } + port := listener.Addr().(*net.TCPAddr).Port + _ = listener.Close() + + pc, err := net.ListenPacket("udp", "127.0.0.1:0") + if err != nil { + t.Fatalf("ListenPacket: %v", err) + } + _ = pc.Close() + + return Config{ + Addr: fmt.Sprintf("127.0.0.1:%d", port), + Net: network, + SinkholeIPv4: netip.MustParseAddr("203.0.113.10"), + SinkholeIPv6: netip.MustParseAddr("2001:db8::10"), + SinkholeDomain: "localhost", + SinkholeTXT: "test", + TTL: 60, + Compress: false, + } +} + +func newAQuery() *dns.Msg { + m := new(dns.Msg) + m.SetQuestion("example.com.", dns.TypeA) + return m +} + +func startDNSService(t *testing.T, conf Config, run *capture.Run) (service.Service, <-chan error) { + t.Helper() + logger := testutil.Logger() + svc := NewService(conf, logger, run) + + errCh := make(chan error, 1) + go func() { errCh <- svc.Start(context.Background()) }() + return svc, errCh +} + +func retryExchange(t *testing.T, client *dns.Client, addr string, m *dns.Msg) (*dns.Msg, time.Duration, error) { + t.Helper() + deadline := time.Now().Add(3 * time.Second) + var lastErr error + var lastRTT time.Duration + for time.Now().Before(deadline) { + resp, rtt, err := client.Exchange(m, addr) + if err == nil && resp != nil { + return resp, rtt, nil + } + lastErr, lastRTT = err, rtt + time.Sleep(20 * time.Millisecond) + } + return nil, lastRTT, lastErr +} + +func discardStartErr(t *testing.T, errCh <-chan error) { + t.Helper() + select { + case err := <-errCh: + if err != nil { + t.Fatalf("service.Start returned error: %v", err) + } + case <-time.After(3 * time.Second): + t.Fatalf("service.Start never returned") + } +} + +// waitDNSCapture waits for the run capture's transport payloads to contain +// want. +func waitDNSCapture(t *testing.T, path, want string) { + t.Helper() + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + if payloads, err := transportPayloads(path); err == nil && strings.Contains(payloads, want) { + return + } + time.Sleep(20 * time.Millisecond) + } + t.Fatalf("run capture %s never contained %q", path, want) +} + +// waitTransportPayloads polls until the transport payloads of the capture +// satisfy cond (tolerating the async flush on connection teardown). +func waitTransportPayloads(t *testing.T, path string, cond func(string) bool) { + t.Helper() + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + joined, err := transportPayloads(path) + if err == nil && cond(joined) { + return + } + time.Sleep(20 * time.Millisecond) + } + joined, _ := transportPayloads(path) + t.Fatalf("capture %s never satisfied predicate (payloads %q)", path, joined) +} + +// transportPayloads concatenates UDP/TCP payloads from a pcapng file. +func transportPayloads(path string) (string, error) { + f, err := os.Open(path) + if err != nil { + return "", err + } + defer f.Close() + r, err := pcapgo.NewNgReader(f, pcapgo.DefaultNgReaderOptions) + if err != nil { + return "", err + } + var sb strings.Builder + for { + data, _, err := r.ReadPacketData() + if err != nil { + break + } + pkt := gopacket.NewPacket(data, layers.LinkTypeEthernet, gopacket.Default) + if u, ok := pkt.Layer(layers.LayerTypeUDP).(*layers.UDP); ok { + sb.Write(u.Payload) + } else if tcp, ok := pkt.Layer(layers.LayerTypeTCP).(*layers.TCP); ok { + sb.Write(tcp.Payload) + } + } + return sb.String(), nil +} diff --git a/internal/dnsserver/config.go b/internal/dnsserver/config.go index 6cdd1b8..09f45ee 100644 --- a/internal/dnsserver/config.go +++ b/internal/dnsserver/config.go @@ -2,14 +2,11 @@ package dnsserver import ( "errors" - "log/slog" "net" "net/netip" - "strings" + "time" - "github.com/miekg/dns" - - "github.com/lachlanharrisdev/gonetsim/internal/service" + "github.com/lachlanharrisdev/gonetsim/internal/netx" ) const AutoIPv4 = "auto" @@ -36,20 +33,6 @@ func AutoSinkholeIPv4() netip.Addr { return netip.MustParseAddr("127.0.0.1") } -func (s *Server) Name() string { - return "DNS" -} - -type Server struct { - conf Config - srvs []*dns.Server - log *slog.Logger -} - -func NewService(conf Config, logger *slog.Logger) service.Service { - return &Server{conf: conf, log: service.NewPrefixedLogger(logger, "DNS")} -} - type Config struct { Addr string Net string @@ -60,21 +43,18 @@ type Config struct { SinkholeTXT string TTL uint32 Compress bool + Capture bool } +// how long to keep a UDP capture writer open after its lastdatagram +const flowIdle = 5 * time.Minute + func (c Config) Validate() error { if c.Addr == "" { return errors.New("listen addr is required") } - if c.Net == "" { - return errors.New("network is required") - } - net := strings.ToLower(strings.TrimSpace(c.Net)) - switch net { - case "udp", "tcp", "both": - // all good my boy - default: - return errors.New("network must be one of: udp, tcp, both") + if err := netx.ValidateNetwork(c.Net, "udp", "tcp", "both"); err != nil { + return err } if !c.SinkholeIPv4.IsValid() { return errors.New("sinkhole ipv4 is required") diff --git a/internal/dnsserver/dns_test.go b/internal/dnsserver/dns_test.go index 8edadb3..67736b5 100644 --- a/internal/dnsserver/dns_test.go +++ b/internal/dnsserver/dns_test.go @@ -2,14 +2,14 @@ package dnsserver import ( "fmt" - "io" - "log/slog" "net" "net/netip" "testing" "time" "github.com/miekg/dns" + + "github.com/lachlanharrisdev/gonetsim/internal/testutil" ) // not a test in of itself; sets up config and server for all record-specific tests (e.g. A, AAAA, TXT) to use, to avoid duplication of setup code in each test @@ -31,7 +31,7 @@ func queryTestsHelper(t *testing.T) (client *dns.Client, addr string, config Con TTL: 60, Compress: false, } - logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + logger := testutil.Logger() srv, err := NewServer(conf, logger) if err != nil { // failed to create server with error @@ -87,7 +87,7 @@ func queryBothTransportsHelper(t *testing.T) (udpClient *dns.Client, tcpClient * TTL: 60, Compress: false, } - logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + logger := testutil.Logger() srvs, err := NewServers(conf, logger) if err != nil { @@ -146,196 +146,118 @@ func TestAutoSinkholeIPv4(t *testing.T) { } } -func TestWildcardDomain(t *testing.T) { - client, addr, config, teardown := queryTestsHelper(t) - defer teardown() - - response := exchange(t, client, addr, "random-beacon-9f3a.malware.example.", dns.TypeA) - if len(response.Answer) != 1 { - t.Fatalf("expected 1 answer, got %d", len(response.Answer)) - } - a, ok := response.Answer[0].(*dns.A) - if !ok { - t.Fatalf("expected *dns.A, got %T", response.Answer[0]) - } - if got := a.A.String(); got != config.SinkholeIPv4.String() { - t.Fatalf("expected %s, got %s", config.SinkholeIPv4.String(), got) +func TestRecordTypes(t *testing.T) { + cases := []struct { + name string + qname string + qtype uint16 + check func(t *testing.T, resp *dns.Msg, conf Config) + }{ + {"wildcard", "random-beacon-9f3a.malware.example.", dns.TypeA, checkA}, + {"A", "example.com.", dns.TypeA, checkA}, + {"AAAA", "example.com.", dns.TypeAAAA, checkAAAA}, + {"TXT", "example.com.", dns.TypeTXT, checkTXT}, + {"CNAME", "example.com.", dns.TypeCNAME, checkDomainTarget}, + {"MX", "example.com.", dns.TypeMX, checkDomainTarget}, + {"NS", "example.com.", dns.TypeNS, checkDomainTarget}, + {"SRV", "_sip._tcp.example.com.", dns.TypeSRV, checkDomainTarget}, + {"PTR", "example.com.", dns.TypePTR, checkDomainTarget}, + {"SOA", "example.com.", dns.TypeSOA, checkSOA}, + {"CAA", "example.com.", dns.TypeCAA, checkCAA}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + client, addr, conf, teardown := queryTestsHelper(t) + defer teardown() + resp := exchange(t, client, addr, tc.qname, tc.qtype) + if len(resp.Answer) != 1 { + t.Fatalf("expected 1 answer, got %d", len(resp.Answer)) + } + tc.check(t, resp, conf) + }) } } -func TestAQuery(t *testing.T) { - client, addr, config, teardown := queryTestsHelper(t) - defer teardown() - - response := exchange(t, client, addr, "example.com.", dns.TypeA) - if len(response.Answer) != 1 { - t.Fatalf("expected 1 answer, got %d", len(response.Answer)) - } - a, ok := response.Answer[0].(*dns.A) +func checkA(t *testing.T, resp *dns.Msg, conf Config) { + t.Helper() + a, ok := resp.Answer[0].(*dns.A) if !ok { - t.Fatalf("expected *dns.A, got %T", response.Answer[0]) + t.Fatalf("expected *dns.A, got %T", resp.Answer[0]) } - if got := a.A.String(); got != config.SinkholeIPv4.String() { - t.Fatalf("expected %s, got %s", config.SinkholeIPv4.String(), got) + if got := a.A.String(); got != conf.SinkholeIPv4.String() { + t.Fatalf("expected %s, got %s", conf.SinkholeIPv4.String(), got) } } -func TestAAAAQuery(t *testing.T) { - client, addr, config, teardown := queryTestsHelper(t) - defer teardown() - - response := exchange(t, client, addr, "example.com.", dns.TypeAAAA) - if len(response.Answer) != 1 { - t.Fatalf("expected 1 answer, got %d", len(response.Answer)) - } - aaaa, ok := response.Answer[0].(*dns.AAAA) +func checkAAAA(t *testing.T, resp *dns.Msg, conf Config) { + t.Helper() + aaaa, ok := resp.Answer[0].(*dns.AAAA) if !ok { - t.Fatalf("expected *dns.AAAA, got %T", response.Answer[0]) + t.Fatalf("expected *dns.AAAA, got %T", resp.Answer[0]) } - if got := aaaa.AAAA.String(); got != config.SinkholeIPv6.String() { - t.Fatalf("expected %s, got %s", config.SinkholeIPv6.String(), got) + if got := aaaa.AAAA.String(); got != conf.SinkholeIPv6.String() { + t.Fatalf("expected %s, got %s", conf.SinkholeIPv6.String(), got) } } -func TestTXTQuery(t *testing.T) { - client, addr, config, teardown := queryTestsHelper(t) - defer teardown() - - response := exchange(t, client, addr, "example.com.", dns.TypeTXT) - if len(response.Answer) != 1 { - t.Fatalf("expected 1 answer, got %d", len(response.Answer)) - } - txt, ok := response.Answer[0].(*dns.TXT) +func checkTXT(t *testing.T, resp *dns.Msg, conf Config) { + t.Helper() + txt, ok := resp.Answer[0].(*dns.TXT) if !ok { - t.Fatalf("expected *dns.TXT, got %T", response.Answer[0]) + t.Fatalf("expected *dns.TXT, got %T", resp.Answer[0]) } if len(txt.Txt) != 1 { t.Fatalf("expected 1 TXT record, got %d", len(txt.Txt)) } - if got := txt.Txt[0]; got != config.SinkholeTXT { - t.Fatalf("expected %s, got %s", config.SinkholeTXT, got) - } -} - -func TestCNAMEQuery(t *testing.T) { - client, addr, config, teardown := queryTestsHelper(t) - defer teardown() - - response := exchange(t, client, addr, "example.com.", dns.TypeCNAME) - if len(response.Answer) != 1 { - t.Fatalf("expected 1 answer, got %d", len(response.Answer)) - } - cname, ok := response.Answer[0].(*dns.CNAME) - if !ok { - t.Fatalf("expected *dns.CNAME, got %T", response.Answer[0]) - } - if got := cname.Target; got != config.SinkholeDomain+"." { - t.Fatalf("expected %s., got %s", config.SinkholeDomain, got) + if got := txt.Txt[0]; got != conf.SinkholeTXT { + t.Fatalf("expected %s, got %s", conf.SinkholeTXT, got) } } -func TestMXQuery(t *testing.T) { - client, addr, config, teardown := queryTestsHelper(t) - defer teardown() - - response := exchange(t, client, addr, "example.com.", dns.TypeMX) - if len(response.Answer) != 1 { - t.Fatalf("expected 1 answer, got %d", len(response.Answer)) - } - mx, ok := response.Answer[0].(*dns.MX) - if !ok { - t.Fatalf("expected *dns.MX, got %T", response.Answer[0]) - } - if got := mx.Mx; got != config.SinkholeDomain+"." { - t.Fatalf("expected %s., got %s", config.SinkholeDomain, got) - } -} - -func TestNSQuery(t *testing.T) { - client, addr, config, teardown := queryTestsHelper(t) - defer teardown() - - response := exchange(t, client, addr, "example.com.", dns.TypeNS) - if len(response.Answer) != 1 { - t.Fatalf("expected 1 answer, got %d", len(response.Answer)) - } - ns, ok := response.Answer[0].(*dns.NS) - if !ok { - t.Fatalf("expected *dns.NS, got %T", response.Answer[0]) - } - if got := ns.Ns; got != config.SinkholeDomain+"." { - t.Fatalf("expected %s., got %s", config.SinkholeDomain, got) - } -} - -func TestSRVQuery(t *testing.T) { - client, addr, config, teardown := queryTestsHelper(t) - defer teardown() - - response := exchange(t, client, addr, "_sip._tcp.example.com.", dns.TypeSRV) - if len(response.Answer) != 1 { - t.Fatalf("expected 1 answer, got %d", len(response.Answer)) - } - srv, ok := response.Answer[0].(*dns.SRV) - if !ok { - t.Fatalf("expected *dns.SRV, got %T", response.Answer[0]) - } - if got := srv.Target; got != config.SinkholeDomain+"." { - t.Fatalf("expected %s., got %s", config.SinkholeDomain, got) - } -} - -func TestPTRQuery(t *testing.T) { - client, addr, config, teardown := queryTestsHelper(t) - defer teardown() - - response := exchange(t, client, addr, "example.com.", dns.TypePTR) - if len(response.Answer) != 1 { - t.Fatalf("expected 1 answer, got %d", len(response.Answer)) - } - ptr, ok := response.Answer[0].(*dns.PTR) - if !ok { - t.Fatalf("expected *dns.PTR, got %T", response.Answer[0]) - } - if got := ptr.Ptr; got != config.SinkholeDomain+"." { - t.Fatalf("expected %s., got %s", config.SinkholeDomain, got) +func checkDomainTarget(t *testing.T, resp *dns.Msg, conf Config) { + t.Helper() + var actual string + switch rr := resp.Answer[0].(type) { + case *dns.CNAME: + actual = rr.Target + case *dns.MX: + actual = rr.Mx + case *dns.NS: + actual = rr.Ns + case *dns.SRV: + actual = rr.Target + case *dns.PTR: + actual = rr.Ptr + default: + t.Fatalf("unexpected type %T", resp.Answer[0]) + } + if want := conf.SinkholeDomain + "."; actual != want { + t.Fatalf("expected %s, got %s", want, actual) } } -func TestSOAQuery(t *testing.T) { - client, addr, config, teardown := queryTestsHelper(t) - defer teardown() - - response := exchange(t, client, addr, "example.com.", dns.TypeSOA) - if len(response.Answer) != 1 { - t.Fatalf("expected 1 answer, got %d", len(response.Answer)) - } - soa, ok := response.Answer[0].(*dns.SOA) +func checkSOA(t *testing.T, resp *dns.Msg, conf Config) { + t.Helper() + soa, ok := resp.Answer[0].(*dns.SOA) if !ok { - t.Fatalf("expected *dns.SOA, got %T", response.Answer[0]) + t.Fatalf("expected *dns.SOA, got %T", resp.Answer[0]) } - if got := soa.Ns; got != config.SinkholeDomain+"." { + if got := soa.Ns; got != conf.SinkholeDomain+"." { t.Fatalf("expected localhost., got %s", got) } - if got := soa.Mbox; got != fmt.Sprintf("hostmaster.%s.", config.SinkholeDomain) { - t.Fatalf("expected hostmaster.%s., got %s", config.SinkholeDomain, got) + if got := soa.Mbox; got != fmt.Sprintf("hostmaster.%s.", conf.SinkholeDomain) { + t.Fatalf("expected hostmaster.%s., got %s", conf.SinkholeDomain, got) } } -func TestCAAQuery(t *testing.T) { - client, addr, config, teardown := queryTestsHelper(t) - defer teardown() - - response := exchange(t, client, addr, "example.com.", dns.TypeCAA) - if len(response.Answer) != 1 { - t.Fatalf("expected 1 answer, god %d", len(response.Answer)) - } - caa, ok := response.Answer[0].(*dns.CAA) +func checkCAA(t *testing.T, resp *dns.Msg, conf Config) { + t.Helper() + caa, ok := resp.Answer[0].(*dns.CAA) if !ok { - t.Fatalf("expected *dns.CAA, got %T", response.Answer[0]) + t.Fatalf("expected *dns.CAA, got %T", resp.Answer[0]) } - if got := caa.Value; got != config.SinkholeDomain { - t.Fatalf("expected %s, got %s", config.SinkholeDomain, got) + if got := caa.Value; got != conf.SinkholeDomain { + t.Fatalf("expected %s, got %s", conf.SinkholeDomain, got) } if got := caa.Tag; got != "issue" { t.Fatalf("expected tag issue, got %s", got) diff --git a/internal/dnsserver/server.go b/internal/dnsserver/server.go index 174b094..1caf6f3 100644 --- a/internal/dnsserver/server.go +++ b/internal/dnsserver/server.go @@ -8,8 +8,31 @@ import ( "strings" "github.com/miekg/dns" + + "github.com/lachlanharrisdev/gonetsim/internal/capture" + "github.com/lachlanharrisdev/gonetsim/internal/netx" + "github.com/lachlanharrisdev/gonetsim/internal/service" ) +type Server struct { + conf Config + srvs []*dns.Server + log *slog.Logger + run *capture.Run + pconns []*capture.PacketConn +} + +func NewService(conf Config, logger *slog.Logger, run *capture.Run) service.Service { + if !conf.Capture { + run = nil + } + return &Server{conf: conf, log: service.NewPrefixedLogger(logger, "DNS"), run: run} +} + +func (s *Server) Name() string { + return "DNS" +} + func NewServers(conf Config, logger *slog.Logger) ([]*dns.Server, error) { h := &handler{ logger: logger, @@ -59,18 +82,37 @@ func (s *Server) Start(ctx context.Context) error { } s.srvs = srvs - netLabel := strings.ToLower(strings.TrimSpace(s.conf.Net)) - if netLabel == "both" { - netLabel = "udp+tcp" + for _, srv := range srvs { + iface, err := s.run.NewInterface("gonetsim dns " + srv.Net) + if err != nil { + return err + } + switch srv.Net { + case "udp": + wrapped, err := netx.ListenUDP(s.conf.Addr, s.run, iface, flowIdle) + if err != nil { + return err + } + srv.PacketConn = wrapped + s.pconns = append(s.pconns, wrapped) + case "tcp": + ln, err := netx.ListenTCP(s.conf.Addr, s.run, iface, nil) + if err != nil { + return err + } + srv.Listener = ln + default: + return fmt.Errorf("unsupported dns network %q", srv.Net) + } } - logger.Info("listening", "on", s.conf.Addr, "net", netLabel, "sinkhole", sinkholeSummary(s.conf)) + logger.Info("listening", "on", s.conf.Addr, "net", netx.DisplayNetwork(s.conf.Net), "sinkhole", sinkholeSummary(s.conf)) errCh := make(chan error, len(srvs)) for _, srv := range srvs { srv := srv go func() { - errCh <- srv.ListenAndServe() + errCh <- srv.ActivateAndServe() }() } @@ -98,6 +140,10 @@ func (s *Server) Stop(ctx context.Context) error { firstErr = err } } + for _, pc := range s.pconns { + pc.CloseAll() + } + s.pconns = nil s.srvs = nil return firstErr } diff --git a/internal/handler/echo.go b/internal/handler/echo.go index 4f94515..93b1c3e 100644 --- a/internal/handler/echo.go +++ b/internal/handler/echo.go @@ -12,7 +12,6 @@ func (EchoHandler) HandleTCP(_ context.Context, conn net.Conn, env Env) error { for { n, err := conn.Read(buf) if n > 0 { - env.Capture.Write("", buf[:n]) if _, werr := conn.Write(buf[:n]); werr != nil { return werr } @@ -24,6 +23,5 @@ func (EchoHandler) HandleTCP(_ context.Context, conn net.Conn, env Env) error { } func (EchoHandler) HandleUDP(_ context.Context, data []byte, _ net.Addr, env Env) ([]byte, error) { - env.Capture.Write("", data) return data, nil } diff --git a/internal/handler/handler.go b/internal/handler/handler.go index aa678f6..9dd5f05 100644 --- a/internal/handler/handler.go +++ b/internal/handler/handler.go @@ -17,16 +17,22 @@ import ( type Env struct { Logger *slog.Logger - Capture *capture.Writer + Capture *capture.Session IdleTimeout time.Duration // connection idle timeout, used by conn:sleep Global *state.Store } -// Handler processes network traffic for a listener. type Handler interface { + TCPHandler + UDPHandler +} + +type TCPHandler interface { // HandleTCP serves a single accepted connection until it is closed. HandleTCP(ctx context.Context, conn net.Conn, env Env) error +} +type UDPHandler interface { // HandleUDP processes a single datagram and returns an optional reply. HandleUDP(ctx context.Context, data []byte, remote net.Addr, env Env) ([]byte, error) } diff --git a/internal/handler/handler_test.go b/internal/handler/handler_test.go index c990a15..4920870 100644 --- a/internal/handler/handler_test.go +++ b/internal/handler/handler_test.go @@ -1,58 +1,29 @@ +////---------------------------------------------------------------------------- +// NOTICE: to save development time, test files (including this) have been +// generated with LLMs. The author(s) do not claim credit for these tests +// and exist purely for maximising code quality and reliability +// +// For more information please see `/.github/AI_USAGE.md` +//----------------------------------------------------------------------------// + package handler import ( "io" "log/slog" "net" - "os" + "net/netip" "path/filepath" "strings" "testing" - "time" - - lua "github.com/yuin/gopher-lua" "github.com/lachlanharrisdev/gonetsim/internal/capture" "github.com/lachlanharrisdev/gonetsim/internal/state" + "github.com/lachlanharrisdev/gonetsim/internal/testutil" ) func discardLogger() *slog.Logger { - return slog.New(slog.NewTextHandler(io.Discard, nil)) -} - -// testCapture opens a capture writer in a temp dir; the returned func reads -// the capture file back. -func testCapture(t *testing.T) (*capture.Writer, func() string) { - t.Helper() - base := t.TempDir() - store, err := capture.NewStore(base, "test") - if err != nil { - t.Fatalf("NewStore: %v", err) - } - w, err := store.Conn("203.0.113.10:1", time.Now()) - if err != nil { - t.Fatalf("Conn: %v", err) - } - t.Cleanup(func() { _ = w.Close() }) - return w, func() string { - entries, err := os.ReadDir(filepath.Join(base, "test")) - if err != nil || len(entries) != 1 { - return "" - } - data, _ := os.ReadFile(filepath.Join(base, "test", entries[0].Name())) - return string(data) - } -} - -// pipe returns a connected pair, closed on test cleanup. -func pipe(t *testing.T) (client, server net.Conn) { - t.Helper() - client, server = net.Pipe() - t.Cleanup(func() { - _ = client.Close() - _ = server.Close() - }) - return client, server + return testutil.Logger() } // servePipe runs h against one end of a pipe; the other end is returned for @@ -85,64 +56,27 @@ func roundtrip(t *testing.T, client net.Conn, payload, reply string) { func TestBuiltins(t *testing.T) { t.Run("tcp echo", func(t *testing.T) { - w, read := testCapture(t) - client, done := servePipe(t, EchoHandler{}, Env{Logger: discardLogger(), Capture: w}) + client, done := servePipe(t, EchoHandler{}, Env{Logger: discardLogger()}) roundtrip(t, client, "abc", "abc") _ = client.Close() if err := <-done; err != nil { t.Fatalf("HandleTCP: %v", err) } - if got := read(); got != "abc" { - t.Fatalf("capture content %q", got) - } - }) - - t.Run("tcp sink", func(t *testing.T) { - w, read := testCapture(t) - client, done := servePipe(t, SinkHandler{}, Env{Logger: discardLogger(), Capture: w}) - roundtrip(t, client, "secret exfil", "") - _ = client.Close() - if err := <-done; err != nil { - t.Fatalf("HandleTCP: %v", err) - } - if got := read(); got != "secret exfil" { - t.Fatalf("capture content %q", got) - } }) - addr, _ := net.ResolveUDPAddr("udp", "203.0.113.10:53") t.Run("udp echo", func(t *testing.T) { + addr, _ := net.ResolveUDPAddr("udp", "203.0.113.10:53") reply, err := EchoHandler{}.HandleUDP(t.Context(), []byte("query"), addr, Env{Logger: discardLogger()}) if err != nil || string(reply) != "query" { t.Fatalf("udp echo: %v %q", err, reply) } }) - - t.Run("udp sink", func(t *testing.T) { - reply, err := SinkHandler{}.HandleUDP(t.Context(), []byte("x"), nil, Env{Logger: discardLogger()}) - if err != nil || reply != nil { - t.Fatalf("udp sink: %v %q", err, reply) - } - }) } func TestNewSpecErrors(t *testing.T) { - cases := []struct { - spec string - baseDir string - }{ - {"", "testdata"}, - {"noscheme", "testdata"}, - {"builtin:nope", "testdata"}, - {"python:foo.py", "testdata"}, - {"lua:missing.lua", "testdata"}, - {"lua:bad_syntax.lua", "testdata"}, - {"lua:no_entry.lua", "testdata"}, - {"lua:sandbox_escape.lua", "testdata"}, - } - for _, tc := range cases { - if _, err := New(tc.spec, tc.baseDir, nil); err == nil { - t.Errorf("New(%q): expected error", tc.spec) + for _, spec := range []string{"", "noscheme", "builtin:nope", "python:foo.py", "lua:missing.lua", "lua:bad_syntax.lua", "lua:no_entry.lua", "lua:sandbox_escape.lua"} { + if _, err := New(spec, "testdata", nil); err == nil { + t.Errorf("New(%q): expected error", spec) } } } @@ -161,28 +95,6 @@ func TestLuaHandler(t *testing.T) { } }) - t.Run("tcp read(n)", func(t *testing.T) { - h, err := NewLua("testdata/read_n.lua", nil) - if err != nil { - t.Fatalf("NewLua: %v", err) - } - client, done := servePipe(t, h, Env{Logger: discardLogger()}) - roundtrip(t, client, "ABCDrest", "got:ABCD") - _ = client.Close() - <-done - }) - - t.Run("tcp read_until headers", func(t *testing.T) { - h, err := NewLua("testdata/read_until.lua", nil) - if err != nil { - t.Fatalf("NewLua: %v", err) - } - client, done := servePipe(t, h, Env{Logger: discardLogger()}) - roundtrip(t, client, "GET / HTTP/1.1\r\nHost: x\r\n\r\nrest", "len:27") - _ = client.Close() - <-done - }) - t.Run("udp packets", func(t *testing.T) { remote, _ := net.ResolveUDPAddr("udp", "203.0.113.10:53531") h, err := NewLua("testdata/packet.lua", nil) @@ -199,135 +111,85 @@ func TestLuaHandler(t *testing.T) { } }) - t.Run("capture and log", func(t *testing.T) { - h, err := NewLua("testdata/capture.lua", nil) + t.Run("capture comment", func(t *testing.T) { + h, err := NewLua("testdata/comment.lua", nil) if err != nil { t.Fatalf("NewLua: %v", err) } - w, read := testCapture(t) - client, done := servePipe(t, h, Env{Logger: discardLogger(), Capture: w}) - roundtrip(t, client, "payload", "") - _ = client.Close() - if err := <-done; err != nil { - t.Fatalf("HandleTCP: %v", err) - } - if got := read(); got != "=== section ===\npayload\n" { - t.Fatalf("capture content %q", got) - } - }) + c1, c2 := net.Pipe() + defer c2.Close() + done := make(chan error, 1) + go func() { + done <- h.HandleTCP(t.Context(), c2, Env{Logger: discardLogger()}) + }() - t.Run("sandbox globals", func(t *testing.T) { - h, err := NewLua("testdata/sandbox_report.lua", nil) + run, err := capture.NewRun(filepath.Join(t.TempDir(), "run.pcapng")) if err != nil { - t.Fatalf("NewLua: %v", err) + t.Fatalf("NewRun: %v", err) } - client, done := servePipe(t, h, Env{Logger: discardLogger()}) - buf := make([]byte, 1024) - n, err := client.Read(buf) + defer func() { _ = run.Close() }() + iface, err := run.NewInterface("test") if err != nil { - t.Fatalf("Read: %v", err) + t.Fatalf("NewInterface: %v", err) } - reply := string(buf[:n]) - if !strings.Contains(reply, "io=nil") || !strings.Contains(reply, "os=nil") || !strings.Contains(reply, "require=nil") { - t.Fatalf("sandbox globals leaked: %q", reply) + ses, err := run.NewSession("tcp", + netip.MustParseAddrPort("127.0.0.1:9"), netip.MustParseAddrPort("203.0.113.10:1"), iface) + if err != nil { + t.Fatalf("NewSession: %v", err) + } + ses.Comment("client said hello") + _ = ses.Write([]byte("hello"), true) + if err := ses.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + _, _ = c1.Write([]byte("hello")) + if err := c1.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if err := <-done; err != nil { + t.Fatalf("HandleTCP: %v", err) } - _ = client.Close() - <-done }) } -func TestLuaStateScopes(t *testing.T) { - cases := []struct { - name string - budget *state.Budget - firstReply string - secondReply string - }{ - {"persistence across connections", state.NewBudget(state.DefaultTotalLimit), "1|conn|yes", "2|conn|yes"}, - {"set failure is graceful", state.NewBudget(3), "1|nil|nil", "2|nil|nil"}, +func TestLuaState(t *testing.T) { + h, err := NewLua("testdata/state.lua", state.NewBudget(state.DefaultTotalLimit)) + if err != nil { + t.Fatalf("NewLua: %v", err) } - for _, tc := range cases { - t.Run(tc.name, func(t *testing.T) { - h, err := NewLua("testdata/state.lua", tc.budget) - if err != nil { - t.Fatalf("NewLua: %v", err) - } - - env := Env{Logger: discardLogger(), Global: state.NewStore(tc.budget)} - - client, done := servePipe(t, h, env) - buf := make([]byte, len(tc.firstReply)) - if _, err := io.ReadFull(client, buf); err != nil { - t.Fatalf("ReadFull: %v", err) - } - if string(buf) != tc.firstReply { - t.Fatalf("first connection = %q, want %q", buf, tc.firstReply) - } - _ = client.Close() - if err := <-done; err != nil { - t.Fatalf("HandleTCP: %v", err) - } - - client, done = servePipe(t, h, env) - buf = make([]byte, len(tc.secondReply)) - if _, err := io.ReadFull(client, buf); err != nil { - t.Fatalf("ReadFull: %v", err) - } - if string(buf) != tc.secondReply { - t.Fatalf("second connection = %q, want %q", buf, tc.secondReply) - } - _ = client.Close() - <-done - }) + env := Env{Logger: discardLogger(), Global: state.NewStore(state.NewBudget(state.DefaultTotalLimit))} + for i, want := range []string{"1|conn|yes", "2|conn|yes"} { + client, done := servePipe(t, h, env) + buf := make([]byte, len(want)) + if _, err := io.ReadFull(client, buf); err != nil { + t.Fatalf("ReadFull: %v", err) + } + if string(buf) != want { + t.Fatalf("connection %d = %q, want %q", i, buf, want) + } + _ = client.Close() + if err := <-done; err != nil { + t.Fatalf("HandleTCP: %v", err) + } } } -func TestLuaConnLimits(t *testing.T) { - flood := func(client net.Conn) { - go func() { - buf := make([]byte, 4096) - for { - if _, err := client.Write(buf); err != nil { - return - } - } - }() +func TestSandboxGlobals(t *testing.T) { + h, err := NewLua("testdata/sandbox_report.lua", nil) + if err != nil { + t.Fatalf("NewLua: %v", err) } - - t.Run("read_line cap", func(t *testing.T) { - client, server := pipe(t) - flood(client) - - lc := newLuaConn(server) - _, err := lc.readLine() - if err == nil || !strings.Contains(err.Error(), "line exceeds") { - t.Fatalf("expected line cap error, got: %v", err) - } - }) - - t.Run("read_until cap", func(t *testing.T) { - client, server := pipe(t) - flood(client) - - lc := newLuaConn(server) - _, err := lc.readUntil([]byte("\r\n")) - if err == nil || !strings.Contains(err.Error(), "read exceeds") { - t.Fatalf("expected read cap error, got: %v", err) - } - }) - - t.Run("read_until across chunks", func(t *testing.T) { - client, server := pipe(t) - go func() { - _, _ = client.Write([]byte("HEAD")) - _, _ = client.Write([]byte("ER:X")) - _, _ = client.Write([]byte("\r\n\r\n")) - }() - - lc := newLuaConn(server) - v, err := lc.readUntil([]byte("\r\n\r\n")) - if err != nil || v != lua.LString("HEADER:X\r\n\r\n") { - t.Fatalf("readUntil: %v %q", err, v) - } - }) + client, done := servePipe(t, h, Env{Logger: discardLogger()}) + buf := make([]byte, 1024) + n, err := client.Read(buf) + if err != nil { + t.Fatalf("Read: %v", err) + } + reply := string(buf[:n]) + if !strings.Contains(reply, "io=nil") || !strings.Contains(reply, "os=nil") || !strings.Contains(reply, "require=nil") { + t.Fatalf("sandbox globals leaked: %q", reply) + } + _ = client.Close() + <-done } diff --git a/internal/handler/lua.go b/internal/handler/lua.go index a5b1503..040bbe2 100644 --- a/internal/handler/lua.go +++ b/internal/handler/lua.go @@ -1,24 +1,17 @@ package handler import ( - "bufio" "bytes" "context" - "crypto/tls" - "errors" "fmt" - "io" - "log/slog" "net" "os" "path/filepath" - "strings" "time" lua "github.com/yuin/gopher-lua" "github.com/yuin/gopher-lua/parse" - "github.com/lachlanharrisdev/gonetsim/internal/capture" "github.com/lachlanharrisdev/gonetsim/internal/state" ) @@ -148,344 +141,3 @@ func (h *LuaHandler) newState(env Env) *lua.LState { registerState(L, "handler", h.handlerState) return L } - -func registerState(L *lua.LState, name string, store *state.Store) { - t := L.NewTable() - registerStateMethods(L, t, store) - L.SetGlobal(name, t) -} - -func registerStateMethods(L *lua.LState, t *lua.LTable, store *state.Store) { - L.SetField(t, "get", L.NewFunction(func(L *lua.LState) int { - if v, ok := store.Get(L.CheckString(2)); ok { - L.Push(lua.LString(v)) - } else { - L.Push(lua.LNil) - } - return 1 - })) - L.SetField(t, "set", L.NewFunction(func(L *lua.LState) int { - if err := store.Set(L.CheckString(2), L.CheckString(3)); err != nil { - L.Push(lua.LFalse) - L.Push(lua.LString(err.Error())) - return 2 - } - L.Push(lua.LTrue) - return 1 - })) - L.SetField(t, "has", L.NewFunction(func(L *lua.LState) int { - L.Push(lua.LBool(store.Has(L.CheckString(2)))) - return 1 - })) - L.SetField(t, "delete", L.NewFunction(func(L *lua.LState) int { - store.Delete(L.CheckString(2)) - return 0 - })) -} - -func openLibs(L *lua.LState) { - lua.OpenBase(L) - lua.OpenString(L) - lua.OpenTable(L) - lua.OpenMath(L) - - str := L.GetGlobal("string").(*lua.LTable) - L.SetField(str, "pack", L.NewFunction(luaPack)) - L.SetField(str, "unpack", L.NewFunction(luaUnpack)) - - // base exposes filesystem helpers; drop them - for _, name := range []string{"dofile", "loadfile", "require"} { - L.SetGlobal(name, lua.LNil) - } -} - -func registerLog(L *lua.LState, logger *slog.Logger) { - log := L.NewTable() - for _, e := range []struct { - name string - level slog.Level - }{ - {"info", slog.LevelInfo}, - {"warn", slog.LevelWarn}, - {"error", slog.LevelError}, - } { - fn := L.NewFunction(func(L *lua.LState) int { - logger.Log(context.Background(), e.level, luaStrings(L)) - return 0 - }) - L.SetField(log, e.name, fn) - } - L.SetGlobal("log", log) - - L.SetGlobal("print", L.NewFunction(func(L *lua.LState) int { - logger.Info(luaStrings(L)) - return 0 - })) -} - -func registerCapture(L *lua.LState, w *capture.Writer) { - capture := L.NewTable() - L.SetField(capture, "write", L.NewFunction(func(L *lua.LState) int { - name := L.CheckString(2) - data := L.CheckString(3) - w.Write(name, []byte(data)) - return 0 - })) - L.SetGlobal("capture", capture) -} - -func luaStrings(L *lua.LState) string { - parts := make([]string, L.GetTop()) - for i := range parts { - parts[i] = L.ToString(i + 1) - } - return strings.Join(parts, " ") -} - -type tlsState interface { - ConnectionState() tls.ConnectionState - HandshakeContext(ctx context.Context) error -} - -// luaConn wraps a net.Conn with a shared buffered reader so read and -// read_line never lose buffered data. -type luaConn struct { - net.Conn - br *bufio.Reader -} - -func newLuaConn(conn net.Conn) *luaConn { - return &luaConn{Conn: conn, br: bufio.NewReader(conn)} -} - -func (lc *luaConn) tls() tlsState { - tc, _ := lc.Conn.(tlsState) - return tc -} - -func (lc *luaConn) ConnectionState() tls.ConnectionState { - if tc := lc.tls(); tc != nil { - return tc.ConnectionState() - } - return tls.ConnectionState{} -} - -func (lc *luaConn) HandshakeContext(ctx context.Context) error { - if tc := lc.tls(); tc != nil { - return tc.HandshakeContext(ctx) - } - return nil -} - -func (lc *luaConn) handshake(ctx context.Context) (tls.ConnectionState, bool) { - tc := lc.tls() - if tc == nil { - return tls.ConnectionState{}, false - } - if !tc.ConnectionState().HandshakeComplete { - _ = tc.HandshakeContext(ctx) - } - st := tc.ConnectionState() - if st.Version == 0 { - return tls.ConnectionState{}, false - } - return st, true -} - -func (lc *luaConn) read(n int) (lua.LValue, error) { - buf := make([]byte, n) - nr, err := lc.br.Read(buf) - if nr > 0 { - return lua.LString(buf[:nr]), nil - } - if errors.Is(err, io.EOF) { - return lua.LNil, nil - } - return nil, err -} - -func (lc *luaConn) readLine() (lua.LValue, error) { - var sb strings.Builder - for { - chunk, err := lc.br.ReadSlice('\n') - sb.Write(chunk) - if err == nil { - return lua.LString(sb.String()), nil - } - if errors.Is(err, bufio.ErrBufferFull) { - if sb.Len() > maxReadLen { - return nil, fmt.Errorf("line exceeds %d bytes", maxReadLen) - } - continue - } - if errors.Is(err, io.EOF) { - if sb.Len() > 0 { - return lua.LString(sb.String()), nil - } - return lua.LNil, nil - } - return nil, err - } -} - -func (lc *luaConn) readUntil(delim []byte) (lua.LValue, error) { - var buf []byte - tmp := make([]byte, 4096) - for { - if i := bytes.Index(buf, delim); i >= 0 { - return lua.LString(buf[:i+len(delim)]), nil - } - if len(buf) > maxReadLen { - return nil, fmt.Errorf("read exceeds %d bytes", maxReadLen) - } - n, err := lc.br.Read(tmp) - buf = append(buf, tmp[:n]...) - if err != nil { - if errors.Is(err, io.EOF) { - if len(buf) > 0 { - return lua.LString(buf), nil - } - return lua.LNil, nil - } - return nil, err - } - } -} - -func registerConn(L *lua.LState, lc *luaConn, ctx context.Context, env Env, connState *state.Store) *lua.LTable { - conn := L.NewTable() - - L.SetField(conn, "read", L.NewFunction(func(L *lua.LState) int { - n := L.CheckInt(2) - if n <= 0 { - L.ArgError(2, "read size must be > 0") - return 0 - } - v, err := lc.read(n) - return pushResult(L, v, err) - })) - - L.SetField(conn, "read_line", L.NewFunction(func(L *lua.LState) int { - v, err := lc.readLine() - return pushResult(L, v, err) - })) - - L.SetField(conn, "read_until", L.NewFunction(func(L *lua.LState) int { - delim := L.CheckString(2) - if delim == "" { - L.ArgError(2, "delimiter must not be empty") - return 0 - } - v, err := lc.readUntil([]byte(delim)) - return pushResult(L, v, err) - })) - - L.SetField(conn, "write", L.NewFunction(func(L *lua.LState) int { - if _, err := lc.Write([]byte(L.CheckString(2))); err != nil { - L.RaiseError("write: %v", err) - } - return 0 - })) - - L.SetField(conn, "sleep", L.NewFunction(func(L *lua.LState) int { - ms := L.CheckInt(2) - if ms < 0 { - L.ArgError(2, "sleep duration must be >= 0") - return 0 - } - d := time.Duration(ms) * time.Millisecond - if d > maxSleep { - L.ArgError(2, "sleep duration exceeds "+maxSleep.String()) - return 0 - } - select { - case <-time.After(d): - case <-ctx.Done(): - L.RaiseError("interrupted") - return 0 - } - // a sleep is script activity, not client inactivity - if env.IdleTimeout > 0 { - _ = lc.SetDeadline(time.Now().Add(env.IdleTimeout)) - } - return 0 - })) - - L.SetField(conn, "close", L.NewFunction(func(L *lua.LState) int { - _ = lc.Close() - return 0 - })) - - L.SetField(conn, "remote", L.NewFunction(func(L *lua.LState) int { - L.Push(lua.LString(lc.RemoteAddr().String())) - return 1 - })) - - L.SetField(conn, "local", L.NewFunction(func(L *lua.LState) int { - L.Push(lua.LString(lc.LocalAddr().String())) - return 1 - })) - - L.SetField(conn, "remote_ip", L.NewFunction(func(L *lua.LState) int { - if tcp, ok := lc.RemoteAddr().(*net.TCPAddr); ok { - L.Push(lua.LString(tcp.IP.String())) - } else { - L.Push(lua.LNil) - } - return 1 - })) - - L.SetField(conn, "remote_port", L.NewFunction(func(L *lua.LState) int { - if tcp, ok := lc.RemoteAddr().(*net.TCPAddr); ok { - L.Push(lua.LNumber(tcp.Port)) - } else { - L.Push(lua.LNil) - } - return 1 - })) - - L.SetField(conn, "local_port", L.NewFunction(func(L *lua.LState) int { - if tcp, ok := lc.LocalAddr().(*net.TCPAddr); ok { - L.Push(lua.LNumber(tcp.Port)) - } else { - L.Push(lua.LNil) - } - return 1 - })) - - L.SetField(conn, "sni", L.NewFunction(func(L *lua.LState) int { - if st, ok := lc.handshake(ctx); ok && st.ServerName != "" { - L.Push(lua.LString(st.ServerName)) - } else { - L.Push(lua.LNil) - } - return 1 - })) - - L.SetField(conn, "tls", L.NewFunction(func(L *lua.LState) int { - st, ok := lc.handshake(ctx) - if !ok { - L.Push(lua.LNil) - return 1 - } - info := L.NewTable() - L.SetField(info, "version", lua.LString(tls.VersionName(st.Version))) - L.SetField(info, "cipher", lua.LString(tls.CipherSuiteName(st.CipherSuite))) - L.Push(info) - return 1 - })) - - registerStateMethods(L, conn, connState) - - return conn -} - -// pushResult returns a value, nil on clean EOF, or raises on failure. -func pushResult(L *lua.LState, v lua.LValue, err error) int { - if err != nil { - L.RaiseError("%v", err) - return 0 - } - L.Push(v) - return 1 -} diff --git a/internal/handler/luabindings.go b/internal/handler/luabindings.go new file mode 100644 index 0000000..7020b11 --- /dev/null +++ b/internal/handler/luabindings.go @@ -0,0 +1,243 @@ +package handler + +import ( + "context" + "crypto/tls" + "log/slog" + "net" + "strings" + "time" + + lua "github.com/yuin/gopher-lua" + + "github.com/lachlanharrisdev/gonetsim/internal/capture" + "github.com/lachlanharrisdev/gonetsim/internal/state" +) + +func registerState(L *lua.LState, name string, store *state.Store) { + t := L.NewTable() + registerStateMethods(L, t, store) + L.SetGlobal(name, t) +} + +func registerStateMethods(L *lua.LState, t *lua.LTable, store *state.Store) { + L.SetField(t, "get", L.NewFunction(func(L *lua.LState) int { + if v, ok := store.Get(L.CheckString(2)); ok { + L.Push(lua.LString(v)) + } else { + L.Push(lua.LNil) + } + return 1 + })) + L.SetField(t, "set", L.NewFunction(func(L *lua.LState) int { + if err := store.Set(L.CheckString(2), L.CheckString(3)); err != nil { + L.Push(lua.LFalse) + L.Push(lua.LString(err.Error())) + return 2 + } + L.Push(lua.LTrue) + return 1 + })) + L.SetField(t, "has", L.NewFunction(func(L *lua.LState) int { + L.Push(lua.LBool(store.Has(L.CheckString(2)))) + return 1 + })) + L.SetField(t, "delete", L.NewFunction(func(L *lua.LState) int { + store.Delete(L.CheckString(2)) + return 0 + })) +} + +func openLibs(L *lua.LState) { + lua.OpenBase(L) + lua.OpenString(L) + lua.OpenTable(L) + lua.OpenMath(L) + + str := L.GetGlobal("string").(*lua.LTable) + L.SetField(str, "pack", L.NewFunction(luaPack)) + L.SetField(str, "unpack", L.NewFunction(luaUnpack)) + + // base exposes filesystem helpers; drop them + for _, name := range []string{"dofile", "loadfile", "require"} { + L.SetGlobal(name, lua.LNil) + } +} + +func registerLog(L *lua.LState, logger *slog.Logger) { + log := L.NewTable() + for _, e := range []struct { + name string + level slog.Level + }{ + {"info", slog.LevelInfo}, + {"warn", slog.LevelWarn}, + {"error", slog.LevelError}, + } { + fn := L.NewFunction(func(L *lua.LState) int { + logger.Log(context.Background(), e.level, luaStrings(L)) + return 0 + }) + L.SetField(log, e.name, fn) + } + L.SetGlobal("log", log) + + L.SetGlobal("print", L.NewFunction(func(L *lua.LState) int { + logger.Info(luaStrings(L)) + return 0 + })) +} + +func registerCapture(L *lua.LState, ses *capture.Session) { + capture := L.NewTable() + L.SetField(capture, "comment", L.NewFunction(func(L *lua.LState) int { + ses.Comment(L.CheckString(2)) + return 0 + })) + L.SetGlobal("capture", capture) +} + +func luaStrings(L *lua.LState) string { + parts := make([]string, L.GetTop()) + for i := range parts { + parts[i] = L.ToString(i + 1) + } + return strings.Join(parts, " ") +} + +func registerConn(L *lua.LState, lc *luaConn, ctx context.Context, env Env, connState *state.Store) *lua.LTable { + conn := L.NewTable() + + L.SetField(conn, "read", L.NewFunction(func(L *lua.LState) int { + n := L.CheckInt(2) + if n <= 0 { + L.ArgError(2, "read size must be > 0") + return 0 + } + v, err := lc.read(n) + return pushResult(L, v, err) + })) + + L.SetField(conn, "read_line", L.NewFunction(func(L *lua.LState) int { + v, err := lc.readLine() + return pushResult(L, v, err) + })) + + L.SetField(conn, "read_until", L.NewFunction(func(L *lua.LState) int { + delim := L.CheckString(2) + if delim == "" { + L.ArgError(2, "delimiter must not be empty") + return 0 + } + v, err := lc.readUntil([]byte(delim)) + return pushResult(L, v, err) + })) + + L.SetField(conn, "write", L.NewFunction(func(L *lua.LState) int { + if _, err := lc.Write([]byte(L.CheckString(2))); err != nil { + L.RaiseError("write: %v", err) + } + return 0 + })) + + L.SetField(conn, "sleep", L.NewFunction(func(L *lua.LState) int { + ms := L.CheckInt(2) + if ms < 0 { + L.ArgError(2, "sleep duration must be >= 0") + return 0 + } + d := time.Duration(ms) * time.Millisecond + if d > maxSleep { + L.ArgError(2, "sleep duration exceeds "+maxSleep.String()) + return 0 + } + select { + case <-time.After(d): + case <-ctx.Done(): + L.RaiseError("interrupted") + return 0 + } + // a sleep is script activity, not client inactivity + if env.IdleTimeout > 0 { + _ = lc.SetDeadline(time.Now().Add(env.IdleTimeout)) + } + return 0 + })) + + L.SetField(conn, "close", L.NewFunction(func(L *lua.LState) int { + _ = lc.Close() + return 0 + })) + + L.SetField(conn, "remote", L.NewFunction(func(L *lua.LState) int { + L.Push(lua.LString(lc.RemoteAddr().String())) + return 1 + })) + + L.SetField(conn, "local", L.NewFunction(func(L *lua.LState) int { + L.Push(lua.LString(lc.LocalAddr().String())) + return 1 + })) + + L.SetField(conn, "remote_ip", L.NewFunction(func(L *lua.LState) int { + if tcp, ok := lc.RemoteAddr().(*net.TCPAddr); ok { + L.Push(lua.LString(tcp.IP.String())) + } else { + L.Push(lua.LNil) + } + return 1 + })) + + L.SetField(conn, "remote_port", L.NewFunction(func(L *lua.LState) int { + if tcp, ok := lc.RemoteAddr().(*net.TCPAddr); ok { + L.Push(lua.LNumber(tcp.Port)) + } else { + L.Push(lua.LNil) + } + return 1 + })) + + L.SetField(conn, "local_port", L.NewFunction(func(L *lua.LState) int { + if tcp, ok := lc.LocalAddr().(*net.TCPAddr); ok { + L.Push(lua.LNumber(tcp.Port)) + } else { + L.Push(lua.LNil) + } + return 1 + })) + + L.SetField(conn, "sni", L.NewFunction(func(L *lua.LState) int { + if st, ok := lc.handshake(ctx); ok && st.ServerName != "" { + L.Push(lua.LString(st.ServerName)) + } else { + L.Push(lua.LNil) + } + return 1 + })) + + L.SetField(conn, "tls", L.NewFunction(func(L *lua.LState) int { + st, ok := lc.handshake(ctx) + if !ok { + L.Push(lua.LNil) + return 1 + } + info := L.NewTable() + L.SetField(info, "version", lua.LString(tls.VersionName(st.Version))) + L.SetField(info, "cipher", lua.LString(tls.CipherSuiteName(st.CipherSuite))) + L.Push(info) + return 1 + })) + + registerStateMethods(L, conn, connState) + + return conn +} + +func pushResult(L *lua.LState, v lua.LValue, err error) int { + if err != nil { + L.RaiseError("%v", err) + return 0 + } + L.Push(v) + return 1 +} diff --git a/internal/handler/luaconn.go b/internal/handler/luaconn.go new file mode 100644 index 0000000..1fdeb57 --- /dev/null +++ b/internal/handler/luaconn.go @@ -0,0 +1,123 @@ +package handler + +import ( + "bufio" + "bytes" + "context" + "crypto/tls" + "errors" + "fmt" + "io" + "net" + "strings" + + lua "github.com/yuin/gopher-lua" +) + +type tlsState interface { + ConnectionState() tls.ConnectionState + HandshakeContext(ctx context.Context) error +} + +type luaConn struct { + net.Conn + br *bufio.Reader +} + +func newLuaConn(conn net.Conn) *luaConn { + return &luaConn{Conn: conn, br: bufio.NewReader(conn)} +} + +func (lc *luaConn) tls() tlsState { + tc, _ := lc.Conn.(tlsState) + return tc +} + +func (lc *luaConn) ConnectionState() tls.ConnectionState { + if tc := lc.tls(); tc != nil { + return tc.ConnectionState() + } + return tls.ConnectionState{} +} + +func (lc *luaConn) HandshakeContext(ctx context.Context) error { + if tc := lc.tls(); tc != nil { + return tc.HandshakeContext(ctx) + } + return nil +} + +func (lc *luaConn) handshake(ctx context.Context) (tls.ConnectionState, bool) { + tc := lc.tls() + if tc == nil { + return tls.ConnectionState{}, false + } + if !tc.ConnectionState().HandshakeComplete { + _ = tc.HandshakeContext(ctx) + } + st := tc.ConnectionState() + if st.Version == 0 { + return tls.ConnectionState{}, false + } + return st, true +} + +func (lc *luaConn) read(n int) (lua.LValue, error) { + buf := make([]byte, n) + nr, err := lc.br.Read(buf) + if nr > 0 { + return lua.LString(buf[:nr]), nil + } + if errors.Is(err, io.EOF) { + return lua.LNil, nil + } + return nil, err +} + +func (lc *luaConn) readLine() (lua.LValue, error) { + var sb strings.Builder + for { + chunk, err := lc.br.ReadSlice('\n') + sb.Write(chunk) + if err == nil { + return lua.LString(sb.String()), nil + } + if errors.Is(err, bufio.ErrBufferFull) { + if sb.Len() > maxReadLen { + return nil, fmt.Errorf("line exceeds %d bytes", maxReadLen) + } + continue + } + if errors.Is(err, io.EOF) { + if sb.Len() > 0 { + return lua.LString(sb.String()), nil + } + return lua.LNil, nil + } + return nil, err + } +} + +func (lc *luaConn) readUntil(delim []byte) (lua.LValue, error) { + var buf []byte + tmp := make([]byte, 4096) + for { + if i := bytes.Index(buf, delim); i >= 0 { + return lua.LString(buf[:i+len(delim)]), nil + } + if len(buf) > maxReadLen { + return nil, fmt.Errorf("read exceeds %d bytes", maxReadLen) + } + n, err := lc.br.Read(tmp) + buf = append(buf, tmp[:n]...) + if err != nil { + if errors.Is(err, io.EOF) { + if len(buf) > 0 { + return lua.LString(buf), nil + } + return lua.LNil, nil + } + return nil, err + } + } +} diff --git a/internal/handler/luapack_test.go b/internal/handler/luapack_test.go deleted file mode 100644 index 8c68f88..0000000 --- a/internal/handler/luapack_test.go +++ /dev/null @@ -1,170 +0,0 @@ -package handler - -import ( - "strings" - "testing" - - lua "github.com/yuin/gopher-lua" -) - -// runLuaN evaluates src through the sandboxed libraries and returns its nret -// return values. -func runLuaN(t *testing.T, src string, nret int) []lua.LValue { - t.Helper() - L := lua.NewState(lua.Options{SkipOpenLibs: true}) - defer L.Close() - openLibs(L) - proto, err := L.LoadString(src) - if err != nil { - t.Fatalf("load %q: %v", src, err) - } - if err := L.CallByParam(lua.P{Fn: proto, NRet: nret, Protect: true}); err != nil { - t.Fatalf("lua %q: %v", src, err) - } - // results land at the top of the stack; gopher-lua leaves evaluation - // leftovers beneath them - vals := make([]lua.LValue, nret) - for i := range vals { - vals[i] = L.Get(-nret + i) - } - L.Pop(nret) - return vals -} - -func lvString(v lua.LValue) string { return string(v.(lua.LString)) } - -func lvNumber(t *testing.T, v lua.LValue) float64 { - t.Helper() - n, ok := v.(lua.LNumber) - if !ok { - t.Fatalf("expected number, got %s (%v)", v.Type().String(), v) - } - return float64(n) -} - -func TestLuaPack(t *testing.T) { - cases := []struct { - name string - src string - nret int - check func(t *testing.T, vals []lua.LValue) - }{ - {"big endian int", `return string.pack(">i4", 1000)`, 1, func(t *testing.T, v []lua.LValue) { - if got := lvString(v[0]); got != "\x00\x00\x03\xe8" { - t.Errorf("got %q", got) - } - }}, - {"little endian int", `return string.pack("i4", 1000)`, 1, func(t *testing.T, v []lua.LValue) { - if got := lvString(v[0]); got != "\xe8\x03\x00\x00" { - t.Errorf("got %q", got) - } - }}, - {"signed sizes", `return string.pack("H", 65535)`, 4, func(t *testing.T, v []lua.LValue) { - want := []string{"\xfe\xff", "\xff", "\xff", "\xff\xff"} - for i, w := range want { - if got := lvString(v[i]); got != w { - t.Errorf("case %d = %q, want %q", i, got, w) - } - } - }}, - {"default sizes", `return #string.pack("i", 1), #string.pack("j", 1), #string.pack("f", 1), #string.pack("d", 1), #string.pack("s", "hi")`, 5, func(t *testing.T, v []lua.LValue) { - want := []float64{4, 8, 4, 8, 10} // plain "s" prefixes an 8-byte length - for i, w := range want { - if got := lvNumber(t, v[i]); got != w { - t.Errorf("size %d = %v, want %v", i, got, w) - } - } - }}, - {"strings", `return string.pack("z", "hi"), string.pack("c4", "ab"), string.pack("d", 1.5)`, 1, func(t *testing.T, v []lua.LValue) { - if got := lvString(v[0]); got != "\x3f\xf8\x00\x00\x00\x00\x00\x00" { - t.Errorf("got %q", got) - } - }}, - {"unpack signed", `return string.unpack("b", "\255"), string.unpack("I2", "\1\2\3\4", 3)`, 2, func(t *testing.T, v []lua.LValue) { - if got := lvNumber(t, v[0]); got != 0x0304 { - t.Errorf("value = %v", got) - } - if got := lvNumber(t, v[1]); got != 5 { - t.Errorf("position = %v", got) - } - }}, - {"unpack length-prefixed", ` - local d = string.pack(">s2", "payload") - local s, pos = string.unpack(">s2", d) - return s, pos, #d - `, 3, func(t *testing.T, v []lua.LValue) { - if got := lvString(v[0]); got != "payload" { - t.Errorf("string = %q", got) - } - if got := lvNumber(t, v[1]); got != 10 { - t.Errorf("position = %v, want 10", got) - } - if got := lvNumber(t, v[2]); got != 9 { - t.Errorf("total = %v, want 9", got) - } - }}, - {"unpack float roundtrip", `local ok, err = pcall(function() return string.unpack(">f", string.pack(">f", 0.5)) end); return ok, err`, 2, func(t *testing.T, v []lua.LValue) { - if v[0] != lua.LTrue { - t.Fatalf("pcall failed: %v", v[1]) - } - if got := lvNumber(t, v[1]); got != 0.5 { - t.Errorf("float roundtrip = %v", got) - } - }}, - } - for _, tc := range cases { - t.Run(tc.name, func(t *testing.T) { - tc.check(t, runLuaN(t, tc.src, tc.nret)) - }) - } -} - -func TestLuaPackErrors(t *testing.T) { - cases := []struct { - src string - wantIn string - }{ - {`local ok, err = pcall(string.pack, "B", 256); return ok, err`, "integer overflow"}, - {`local ok, err = pcall(string.pack, "i2", 70000); return ok, err`, "integer overflow"}, - {`local ok, err = pcall(string.pack, "i", 1.5); return ok, err`, "no integer representation"}, - {`local ok, err = pcall(string.pack, "z", "a\0b"); return ok, err`, "string contains zeros"}, - {`local ok, err = pcall(string.pack, "c2", "abc"); return ok, err`, "string longer than given size"}, - {`local ok, err = pcall(string.pack, "c", "x"); return ok, err`, "missing size for format option 'c'"}, - {`local ok, err = pcall(string.pack, "!", "x"); return ok, err`, "not supported"}, - {`local ok, err = pcall(string.pack, "q", 1); return ok, err`, "invalid format option"}, - {`local ok, err = pcall(string.pack, "i9", 1); return ok, err`, "out of limits"}, - {`local ok, err = pcall(string.unpack, ">i4", "\1\2"); return ok, err`, "data string too short"}, - {`local ok, err = pcall(string.unpack, "z", "no terminator"); return ok, err`, "zero terminator"}, - {`local ok, err = pcall(string.unpack, ">s2", "\0\5ab"); return ok, err`, "data string too short"}, - {`local ok, err = pcall(string.unpack, "b", "\1", 5); return ok, err`, "initial position out of string"}, - } - for _, tc := range cases { - vals := runLuaN(t, tc.src, 2) - if vals[0] != lua.LFalse { - t.Errorf("%s: expected pcall failure, got %v", tc.src, vals[0]) - } - if got := lvString(vals[1]); !strings.Contains(got, tc.wantIn) { - t.Errorf("%s: error %q does not contain %q", tc.src, got, tc.wantIn) - } - } -} diff --git a/internal/handler/sink.go b/internal/handler/sink.go index 8295be2..790ed1f 100644 --- a/internal/handler/sink.go +++ b/internal/handler/sink.go @@ -11,10 +11,7 @@ type SinkHandler struct{} func (SinkHandler) HandleTCP(_ context.Context, conn net.Conn, env Env) error { buf := make([]byte, 32*1024) for { - n, err := conn.Read(buf) - if n > 0 { - env.Capture.Write("", buf[:n]) - } + _, err := conn.Read(buf) if err != nil { return readError(err) } @@ -22,6 +19,5 @@ func (SinkHandler) HandleTCP(_ context.Context, conn net.Conn, env Env) error { } func (SinkHandler) HandleUDP(_ context.Context, data []byte, _ net.Addr, env Env) ([]byte, error) { - env.Capture.Write("", data) return nil, nil } diff --git a/internal/handler/sleep_sni_test.go b/internal/handler/sleep_sni_test.go deleted file mode 100644 index 8d0bdf3..0000000 --- a/internal/handler/sleep_sni_test.go +++ /dev/null @@ -1,71 +0,0 @@ -package handler - -import ( - "io" - "strings" - "testing" - "time" -) - -func TestLuaSleepAndSNI(t *testing.T) { - t.Run("sleep resets idle deadline", func(t *testing.T) { - h, err := NewLua("testdata/sleep.lua", nil) - if err != nil { - t.Fatalf("NewLua: %v", err) - } - client, done := servePipe(t, h, Env{Logger: discardLogger(), IdleTimeout: 100 * time.Millisecond}) - - // the read after sleep(200) must still succeed even though more - // than IdleTimeout passed since the first read - go func() { - _, _ = client.Write([]byte("go")) - time.Sleep(50 * time.Millisecond) - _, _ = client.Write([]byte("again")) - }() - - start := time.Now() - buf := make([]byte, len("after-sleep")) - if _, err := io.ReadFull(client, buf); err != nil { - t.Fatalf("ReadFull: %v", err) - } - if string(buf) != "after-sleep" { - t.Fatalf("unexpected reply %q", buf) - } - if time.Since(start) < 150*time.Millisecond { - t.Fatalf("sleep(200) returned early") - } - _ = client.Close() - if err := <-done; err != nil { - t.Fatalf("HandleTCP: %v", err) - } - }) - - t.Run("sleep cap", func(t *testing.T) { - h, err := NewLua("testdata/sleep_cap.lua", nil) - if err != nil { - t.Fatalf("NewLua: %v", err) - } - client, done := servePipe(t, h, Env{Logger: discardLogger()}) - _, _ = client.Write([]byte("go")) - if err := <-done; err == nil || !strings.Contains(err.Error(), "exceeds") { - t.Fatalf("expected sleep cap error, got: %v", err) - } - }) - - t.Run("sni nil on plain conn", func(t *testing.T) { - h, err := NewLua("testdata/sni.lua", nil) - if err != nil { - t.Fatalf("NewLua: %v", err) - } - client, done := servePipe(t, h, Env{Logger: discardLogger()}) - buf := make([]byte, 6) - if _, err := io.ReadFull(client, buf); err != nil { - t.Fatalf("ReadFull: %v", err) - } - if string(buf) != "no-sni" { - t.Fatalf("expected no-sni, got %q", buf) - } - _ = client.Close() - <-done - }) -} diff --git a/internal/handler/testdata/capture.lua b/internal/handler/testdata/capture.lua deleted file mode 100644 index 6942edd..0000000 --- a/internal/handler/testdata/capture.lua +++ /dev/null @@ -1,6 +0,0 @@ --- Captures whatever the client sends under a named section. -function handle(conn) - local data = conn:read(1024) - capture:write("section", data) - log:info("captured " .. #data .. " bytes") -end diff --git a/internal/handler/testdata/comment.lua b/internal/handler/testdata/comment.lua new file mode 100644 index 0000000..ad58745 --- /dev/null +++ b/internal/handler/testdata/comment.lua @@ -0,0 +1,5 @@ +-- comment on the next frame the client sends +function handle(conn) + local data = conn:read(1024) + capture:comment("client sent " .. #data .. " bytes") +end \ No newline at end of file diff --git a/internal/httpserver/capture_test.go b/internal/httpserver/capture_test.go new file mode 100644 index 0000000..d036a9f --- /dev/null +++ b/internal/httpserver/capture_test.go @@ -0,0 +1,242 @@ +// //---------------------------------------------------------------------------- +// // NOTICE: to save development time, test files (including this) have been +// // generated with LLMs. The author(s) do not claim credit for these tests +// // and exist purely for maximising code quality and reliability +// // +// // For more information please see `/.github/AI_USAGE.md` +// //----------------------------------------------------------------------------// + +package httpserver + +import ( + "context" + "crypto/tls" + "io" + "net" + "net/http" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/google/gopacket" + "github.com/google/gopacket/layers" + "github.com/google/gopacket/pcapgo" + + "github.com/lachlanharrisdev/gonetsim/internal/capture" + "github.com/lachlanharrisdev/gonetsim/internal/service" + "github.com/lachlanharrisdev/gonetsim/internal/testutil" + "github.com/lachlanharrisdev/gonetsim/internal/tlsprovider" +) + +func testRun(t *testing.T) (*capture.Run, string) { + t.Helper() + path := filepath.Join(t.TempDir(), "run.pcapng") + run, err := capture.NewRun(path) + if err != nil { + t.Fatalf("NewRun: %v", err) + } + t.Cleanup(func() { _ = run.Close() }) + return run, path +} + +func TestService_CapturesHTTP(t *testing.T) { + conf := Config{ + Addr: freeTCPAddr(t), + StatusCode: http.StatusOK, + Mode: "fake", + Capture: true, + } + run, path := testRun(t) + svc, errCh := startHTTPService(t, conf, run) + + get := func(url string) *http.Response { + _, resp := retryingGet(t, http.DefaultClient, url) + return resp + } + get("http://" + conf.Addr + "/warmup") + resp := get("http://" + conf.Addr + "/hello") + _, _ = io.Copy(io.Discard, resp.Body) + _ = resp.Body.Close() + + waitHTTPCapture(t, path, "GET /hello") + waitTCPPayloadsContain(t, path, "HTTP/1.1 200") + + svc.Stop(context.Background()) //nolint:errcheck,gosec + discardStartErr(t, errCh) +} + +func TestService_CapturesHTTPS(t *testing.T) { + dir := t.TempDir() + certPEM, keyPEM, _, err := tlsprovider.GenerateSelfSignedWithCA(tlsprovider.SelfSignedOptions{DNSNames: []string{"localhost"}}) + if err != nil { + t.Fatalf("GenerateSelfSignedWithCA: %v", err) + } + certPath := filepath.Join(dir, "cert.pem") + keyPath := filepath.Join(dir, "key.pem") + if err := os.WriteFile(certPath, certPEM, 0o644); err != nil { + t.Fatalf("WriteFile: %v", err) + } + if err := os.WriteFile(keyPath, keyPEM, 0o600); err != nil { + t.Fatalf("WriteFile: %v", err) + } + + conf := Config{ + Addr: freeTCPAddr(t), + StatusCode: http.StatusOK, + Mode: "fake", + TLS: &tlsprovider.Config{CertFile: certPath, KeyFile: keyPath}, + Capture: true, + } + run, path := testRun(t) + svc, errCh := startHTTPService(t, conf, run) + + client := &http.Client{ + Timeout: 3 * time.Second, + Transport: &http.Transport{ + TLSClientConfig: &tls.Config{InsecureSkipVerify: true}, + }, + } + _, resp := retryingGet(t, client, "https://localhost:"+portOf(t, conf.Addr)+"/warmup") + _, _ = io.Copy(io.Discard, resp.Body) + _ = resp.Body.Close() + _, resp = retryingGet(t, client, "https://localhost:"+portOf(t, conf.Addr)+"/secure") + _, _ = io.Copy(io.Discard, resp.Body) + _ = resp.Body.Close() + + // TLS capture is ciphertext, so assert the run capture holds TLS + // records in both directions rather than plaintext content. + waitHTTPCapture(t, path, "") + waitTCPPayloads(t, path, func(joined string) bool { + return strings.Contains(joined, "\x16") && strings.Contains(joined, "\x17") + }) + + svc.Stop(context.Background()) //nolint:errcheck,gosec + discardStartErr(t, errCh) +} + +// retryingGet issues a GET with Connection: close (forcing the server to +// close the connection so its capture is flushed), retrying while the service +// is still starting up. +func retryingGet(t *testing.T, client *http.Client, url string) (status int, resp *http.Response) { + t.Helper() + deadline := time.Now().Add(3 * time.Second) + var lastErr error + for time.Now().Before(deadline) { + req, err := http.NewRequest(http.MethodGet, url, nil) + if err != nil { + t.Fatalf("NewRequest: %v", err) + } + req.Header.Set("Connection", "close") + r, err := client.Do(req) + if err == nil { + return r.StatusCode, r + } + lastErr = err + time.Sleep(20 * time.Millisecond) + } + t.Fatalf("GET %s: %v", url, lastErr) + return 0, nil +} + +func startHTTPService(t *testing.T, conf Config, run *capture.Run) (service.Service, <-chan error) { + t.Helper() + logger := testutil.Logger() + svc := NewService(conf, logger, run) + + errCh := make(chan error, 1) + go func() { errCh <- svc.Start(context.Background()) }() + return svc, errCh +} + +func discardStartErr(t *testing.T, errCh <-chan error) { + t.Helper() + select { + case err := <-errCh: + if err != nil { + t.Fatalf("service.Start returned error: %v", err) + } + case <-time.After(3 * time.Second): + t.Fatalf("service.Start never returned") + } +} + +// waitHTTPCapture waits for the run capture's payloads to contain want +// (or for any packets when want is empty). +func waitHTTPCapture(t *testing.T, path, want string) { + t.Helper() + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + if payloads, err := tcpPayloads(path); err == nil && (want == "" || strings.Contains(payloads, want)) { + return + } + time.Sleep(20 * time.Millisecond) + } + t.Fatalf("run capture %s never contained %q", path, want) +} + +func waitTCPPayloadsContain(t *testing.T, path, want string) { + t.Helper() + waitTCPPayloads(t, path, func(joined string) bool { + return strings.Contains(joined, want) + }) +} + +func waitTCPPayloads(t *testing.T, path string, cond func(string) bool) { + t.Helper() + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + joined, err := tcpPayloads(path) + if err == nil && cond(joined) { + return + } + time.Sleep(20 * time.Millisecond) + } + joined, _ := tcpPayloads(path) + t.Fatalf("capture %s never satisfied predicate (payloads %q)", path, joined) +} + +func tcpPayloads(path string) (string, error) { + f, err := os.Open(path) + if err != nil { + return "", err + } + defer f.Close() + r, err := pcapgo.NewNgReader(f, pcapgo.DefaultNgReaderOptions) + if err != nil { + return "", err + } + var sb strings.Builder + for { + data, _, err := r.ReadPacketData() + if err != nil { + break + } + pkt := gopacket.NewPacket(data, layers.LinkTypeEthernet, gopacket.Default) + if t, ok := pkt.Layer(layers.LayerTypeTCP).(*layers.TCP); ok { + sb.Write(t.Payload) + } + } + return sb.String(), nil +} + +func freeTCPAddr(t *testing.T) string { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("Listen: %v", err) + } + addr := ln.Addr().String() + _ = ln.Close() + return addr +} + +func portOf(t *testing.T, addr string) string { + t.Helper() + _, port, err := net.SplitHostPort(addr) + if err != nil { + t.Fatalf("SplitHostPort(%q): %v", addr, err) + } + return port +} diff --git a/internal/httpserver/config.go b/internal/httpserver/config.go index 3cbee73..287eb08 100644 --- a/internal/httpserver/config.go +++ b/internal/httpserver/config.go @@ -3,34 +3,12 @@ package httpserver import ( "errors" "fmt" - "log/slog" - "net/http" "os" - "github.com/lachlanharrisdev/gonetsim/internal/service" + "github.com/lachlanharrisdev/gonetsim/internal/netx" "github.com/lachlanharrisdev/gonetsim/internal/tlsprovider" ) -func (s *Server) Name() string { - return s.name -} - -type Server struct { - name string - conf Config - srv *http.Server - log *slog.Logger -} - -func NewService(conf Config, logger *slog.Logger) service.Service { - name := "HTTP" - if conf.TLS != nil { - name = "HTTPS" - } - - return &Server{name: name, conf: conf.normalize(), log: service.NewPrefixedLogger(logger, name)} -} - type Config struct { Addr string @@ -48,6 +26,9 @@ type Config struct { // root directory to serve files // only used in real mode RootDir string + + // write every connection to the run pcapng file + Capture bool } // normalize fills in defaults that can't be expressed as zero values. @@ -64,8 +45,8 @@ func (c Config) Validate() error { if c.Addr == "" { return errors.New("listen addr is required") } - if c.StatusCode != 0 && (c.StatusCode < 100 || c.StatusCode > 599) { - return fmt.Errorf("status code must be 0 or between 100 and 599, was %d", c.StatusCode) + if err := netx.ValidateStatus(c.StatusCode); err != nil { + return err } if c.TLS != nil { if err := c.TLS.Validate(); err != nil { diff --git a/internal/httpserver/http_test.go b/internal/httpserver/http_test.go index f8b7a7b..757964d 100644 --- a/internal/httpserver/http_test.go +++ b/internal/httpserver/http_test.go @@ -1,8 +1,10 @@ -// -------- +////---------------------------------------------------------------------------- // NOTICE: to save development time, test files (including this) have been // generated with LLMs. The author(s) do not claim credit for these tests // and exist purely for maximising code quality and reliability -// -------- +// +// For more information please see `/.github/AI_USAGE.md` +//----------------------------------------------------------------------------// package httpserver @@ -10,7 +12,6 @@ import ( "context" "crypto/tls" "io" - "log/slog" "net" "net/http" "os" @@ -19,6 +20,7 @@ import ( "testing" "time" + "github.com/lachlanharrisdev/gonetsim/internal/testutil" "github.com/lachlanharrisdev/gonetsim/internal/tlsprovider" ) @@ -32,7 +34,7 @@ func TestHTTPServer_Smoke(t *testing.T) { t.Fatalf("listen: %v", err) } - logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + logger := testutil.Logger() srv, err := NewServer(Config{Addr: "127.0.0.1:0", StatusCode: http.StatusCreated}, nil, logger) if err != nil { // failed to create server with error @@ -99,7 +101,7 @@ func TestHTTPSServer_Smoke(t *testing.T) { t.Fatalf("GenerateSelfSigned: %v", err) } - logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + logger := testutil.Logger() srv, err := NewServer(Config{Addr: "127.0.0.1:0", StatusCode: http.StatusOK}, nil, logger) if err != nil { // failed to create https server with error @@ -211,7 +213,7 @@ func startRealServer(t *testing.T, rootDir string, statusCode int) (*http.Server t.Fatalf("listen: %v", err) } - logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + logger := testutil.Logger() srv, err := NewServer(Config{ Addr: "127.0.0.1:0", Mode: "real", @@ -444,110 +446,49 @@ func TestRealHandler_ConditionalRequestNotOverridden(t *testing.T) { // --- security tests --- -func TestRealHandler_DirectoryTraversalBlocked(t *testing.T) { - // Write a sentinel file one level above the root dir. - parent := t.TempDir() - secret := filepath.Join(parent, "secret.txt") - if err := os.WriteFile(secret, []byte("secret contents"), 0o644); err != nil { - t.Fatalf("WriteFile: %v", err) - } - - // The server root is a subdirectory; secret.txt is outside it. - root := filepath.Join(parent, "www") - if err := os.MkdirAll(root, 0o755); err != nil { - t.Fatalf("MkdirAll: %v", err) - } - - _, base := startRealServer(t, root, 0) - - // Classic traversal attempt - resp := mustGet(t, http.DefaultClient, base+"/../secret.txt") - defer resp.Body.Close() //nolint:errcheck - - // Must not serve the file — 404 or 400 are both acceptable - if resp.StatusCode == http.StatusOK { - body, _ := io.ReadAll(resp.Body) - t.Fatalf("traversal succeeded — got 200 with body: %q", string(body)) - } -} - -func TestRealHandler_EncodedTraversalBlocked(t *testing.T) { - parent := t.TempDir() - secret := filepath.Join(parent, "secret.txt") - if err := os.WriteFile(secret, []byte("secret contents"), 0o644); err != nil { - t.Fatalf("WriteFile: %v", err) - } - - root := filepath.Join(parent, "www") - if err := os.MkdirAll(root, 0o755); err != nil { - t.Fatalf("MkdirAll: %v", err) - } - - _, base := startRealServer(t, root, 0) - - // URL-encoded traversal: %2e%2e = ".." - // http.DefaultClient will usually normalise this, but worth having - resp := mustGet(t, http.DefaultClient, base+"/%2e%2e/secret.txt") - defer resp.Body.Close() //nolint:errcheck - - if resp.StatusCode == http.StatusOK { - body, _ := io.ReadAll(resp.Body) - t.Fatalf("encoded traversal succeeded — got 200 with body: %q", string(body)) - } -} - -func TestRealHandler_MiddlePathTraversalBlocked(t *testing.T) { - parent := t.TempDir() - secret := filepath.Join(parent, "secret.txt") - if err := os.WriteFile(secret, []byte("secret contents"), 0o644); err != nil { - t.Fatalf("WriteFile: %v", err) - } - - root := filepath.Join(parent, "www") - if err := os.MkdirAll(root, 0o755); err != nil { - t.Fatalf("MkdirAll: %v", err) - } - - _, base := startRealServer(t, root, 0) - - // .. in the middle of the path must still be contained within root. - resp := mustGet(t, http.DefaultClient, base+"/a/../secret.txt") - defer resp.Body.Close() //nolint:errcheck - - if resp.StatusCode == http.StatusOK { - body, _ := io.ReadAll(resp.Body) - t.Fatalf("middle traversal succeeded — got 200 with body: %q", string(body)) - } -} - -func TestRealHandler_PlainTraversalPath(t *testing.T) { - parent := t.TempDir() - secret := filepath.Join(parent, "secret.txt") - if err := os.WriteFile(secret, []byte("secret contents"), 0o644); err != nil { - t.Fatalf("WriteFile: %v", err) - } - - root := filepath.Join(parent, "www") - if err := os.MkdirAll(root, 0o755); err != nil { - t.Fatalf("MkdirAll: %v", err) - } - - _, base := startRealServer(t, root, 0) - - // Raw .. components that survive URL parsing. - resp := mustGet(t, http.DefaultClient, base+"/../../secret.txt") - defer resp.Body.Close() //nolint:errcheck - - if resp.StatusCode == http.StatusOK { - body, _ := io.ReadAll(resp.Body) - t.Fatalf("raw traversal succeeded — got 200 with body: %q", string(body)) +func TestRealHandler_TraversalBlocked(t *testing.T) { + paths := []struct { + name string + path string + }{ + {"classic", "/../secret.txt"}, + {"encoded", "/%2e%2e/secret.txt"}, + {"middle", "/a/../secret.txt"}, + {"raw", "/../../secret.txt"}, + } + for _, tc := range paths { + t.Run(tc.name, func(t *testing.T) { + // Write a sentinel file one level above the root dir. + parent := t.TempDir() + secret := filepath.Join(parent, "secret.txt") + if err := os.WriteFile(secret, []byte("secret contents"), 0o644); err != nil { + t.Fatalf("WriteFile: %v", err) + } + + // The server root is a subdirectory; secret.txt is outside it. + root := filepath.Join(parent, "www") + if err := os.MkdirAll(root, 0o755); err != nil { + t.Fatalf("MkdirAll: %v", err) + } + + _, base := startRealServer(t, root, 0) + + resp := mustGet(t, http.DefaultClient, base+tc.path) + defer resp.Body.Close() //nolint:errcheck + + // Must not serve the file — 404 or 400 are both acceptable + if resp.StatusCode == http.StatusOK { + body, _ := io.ReadAll(resp.Body) + t.Fatalf("traversal %q succeeded — got 200 with body: %q", tc.path, string(body)) + } + }) } } // --- config validation tests --- func TestNewServer_RealMode_MissingRootDirReturnsError(t *testing.T) { - logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + logger := testutil.Logger() _, err := NewServer(Config{ Addr: "127.0.0.1:0", Mode: "real", @@ -559,7 +500,7 @@ func TestNewServer_RealMode_MissingRootDirReturnsError(t *testing.T) { } func TestNewServer_RealMode_NonexistentRootDirReturnsError(t *testing.T) { - logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + logger := testutil.Logger() _, err := NewServer(Config{ Addr: "127.0.0.1:0", Mode: "real", diff --git a/internal/httpserver/server.go b/internal/httpserver/server.go index 5bc98d1..19f8ff8 100644 --- a/internal/httpserver/server.go +++ b/internal/httpserver/server.go @@ -5,11 +5,39 @@ import ( "crypto/tls" "errors" "log/slog" - "net" "net/http" + "strings" "time" + + "github.com/lachlanharrisdev/gonetsim/internal/capture" + "github.com/lachlanharrisdev/gonetsim/internal/netx" + "github.com/lachlanharrisdev/gonetsim/internal/service" ) +type Server struct { + name string + conf Config + srv *http.Server + log *slog.Logger + run *capture.Run +} + +func NewService(conf Config, logger *slog.Logger, run *capture.Run) service.Service { + name := "HTTP" + if conf.TLS != nil { + name = "HTTPS" + } + if !conf.Capture { + run = nil + } + + return &Server{name: name, conf: conf.normalize(), log: service.NewPrefixedLogger(logger, name), run: run} +} + +func (s *Server) Name() string { + return s.name +} + func NewServer(conf Config, handler http.Handler, logger *slog.Logger) (*http.Server, error) { if err := conf.Validate(); err != nil { return nil, err @@ -42,20 +70,23 @@ func (s *Server) Start(ctx context.Context) error { } s.srv = srv - ln, err := net.Listen("tcp", s.conf.Addr) - if err != nil { - return err - } - defer func() { _ = ln.Close() }() - + var tlsConf *tls.Config if s.conf.TLS != nil { - tlsConf, err := s.conf.TLS.TLSConfig() + tlsConf, err = s.conf.TLS.TLSConfig() if err != nil { return err } srv.TLSConfig = tlsConf - ln = tls.NewListener(ln, tlsConf) } + iface, err := s.run.NewInterface("gonetsim " + strings.ToLower(s.name) + " tcp") + if err != nil { + return err + } + ln, err := netx.ListenTCP(s.conf.Addr, s.run, iface, tlsConf) + if err != nil { + return err + } + defer func() { _ = ln.Close() }() logger.Info("listening", "on", s.conf.Addr, "mode", s.conf.Mode) if err := s.srv.Serve(ln); err != nil && !errors.Is(err, http.ErrServerClosed) { diff --git a/internal/listener/config.go b/internal/listener/config.go index 1459fe9..fcb63a9 100644 --- a/internal/listener/config.go +++ b/internal/listener/config.go @@ -3,10 +3,10 @@ package listener import ( "errors" "fmt" - "net" "strings" "time" + "github.com/lachlanharrisdev/gonetsim/internal/netx" "github.com/lachlanharrisdev/gonetsim/internal/tlsprovider" ) @@ -20,28 +20,19 @@ type Config struct { Capture bool // BaseDir is the directory relative handler script paths resolve against. BaseDir string - // CaptureDir overrides the base directory for capture files. - // When empty, capture.DefaultDir is used. - CaptureDir string } func (c Config) Validate() error { if strings.TrimSpace(c.Name) == "" { return errors.New("name is required") } - if c.Network == "" { - return errors.New("network is required") - } - switch c.Network { - case "tcp", "udp": - // ok - default: - return errors.New("network must be one of: tcp, udp") + if err := netx.ValidateNetwork(c.Network, "tcp", "udp"); err != nil { + return err } if c.Addr == "" { return errors.New("listen addr is required") } - if _, err := net.ResolveTCPAddr("tcp", c.Addr); err != nil { + if _, err := netx.ParseAddr(c.Addr); err != nil { return fmt.Errorf("invalid listen addr %q (expected host:port): %w", c.Addr, err) } if strings.TrimSpace(c.HandlerSpec) == "" { diff --git a/internal/listener/listener_test.go b/internal/listener/listener_test.go index 2ab6657..9022425 100644 --- a/internal/listener/listener_test.go +++ b/internal/listener/listener_test.go @@ -1,3 +1,11 @@ +////---------------------------------------------------------------------------- +// NOTICE: to save development time, test files (including this) have been +// generated with LLMs. The author(s) do not claim credit for these tests +// and exist purely for maximising code quality and reliability +// +// For more information please see `/.github/AI_USAGE.md` +//----------------------------------------------------------------------------// + package listener import ( @@ -12,38 +20,35 @@ import ( "testing" "time" + "github.com/google/gopacket" + "github.com/google/gopacket/layers" + "github.com/google/gopacket/pcapgo" "github.com/lachlanharrisdev/gonetsim/internal/capture" "github.com/lachlanharrisdev/gonetsim/internal/service" + "github.com/lachlanharrisdev/gonetsim/internal/testutil" "github.com/lachlanharrisdev/gonetsim/internal/tlsprovider" ) func testLogger() *slog.Logger { - return slog.New(slog.NewTextHandler(io.Discard, nil)) + return testutil.Logger() } -// freePort reserves an ephemeral port, releases it, and returns its address. -func freePort(t *testing.T, network string) string { +func testRun(t *testing.T) (*capture.Run, string) { t.Helper() - if network == "udp" { - pc, err := net.ListenPacket("udp", "127.0.0.1:0") - if err != nil { - t.Fatalf("ListenPacket: %v", err) - } - addr := pc.LocalAddr().String() - _ = pc.Close() - return addr - } - ln, err := net.Listen("tcp", "127.0.0.1:0") + path := filepath.Join(t.TempDir(), "run.pcapng") + run, err := capture.NewRun(path) if err != nil { - t.Fatalf("Listen: %v", err) + t.Fatalf("NewRun: %v", err) } - addr := ln.Addr().String() - _ = ln.Close() - return addr + t.Cleanup(func() { _ = run.Close() }) + return run, path +} + +func freePort(t *testing.T, network string) string { + t.Helper() + return testutil.FreePort(t, network) } -// startService runs svc in the background; cleanup cancels it and waits for -// Start to return. func startService(t *testing.T, svc service.Service) { t.Helper() ctx, cancel := context.WithCancel(context.Background()) @@ -105,46 +110,79 @@ func echoConfig(t *testing.T) Config { HandlerSpec: "builtin:echo", ReadTimeout: 5 * time.Second, Capture: true, - CaptureDir: t.TempDir(), } } -func captureFile(t *testing.T, dir, listener string) string { +// waitTransportFrames polls the capture until the transport payload sequence +// matches want, tolerating the async flush that follows connection teardown. +func waitTransportFrames(t *testing.T, path, proto string, want []string) { t.Helper() - deadline := time.Now().Add(2 * time.Second) + deadline := time.Now().Add(3 * time.Second) for time.Now().Before(deadline) { - entries, err := os.ReadDir(filepath.Join(dir, listener)) - if err == nil && len(entries) > 0 { - data, err := os.ReadFile(filepath.Join(dir, listener, entries[0].Name())) - if err != nil { - t.Fatalf("ReadFile: %v", err) - } - return string(data) + got, err := transportPayloads(path, proto) + if err == nil && strings.Join(got, "|") == strings.Join(want, "|") { + return } time.Sleep(20 * time.Millisecond) } - t.Fatalf("no capture file appeared in %s", dir) - return "" + got, _ := transportPayloads(path, proto) + t.Fatalf("payload sequence never matched %q (proto %s), last saw %q", strings.Join(want, "|"), proto, strings.Join(got, "|")) } -func luaConfig(t *testing.T, name, script string) Config { - return Config{ - Name: name, - Network: "tcp", - Addr: freePort(t, "tcp"), - HandlerSpec: "lua:" + script, - BaseDir: "../handler/testdata", - ReadTimeout: 5 * time.Second, +// waitSubstringFrames polls until the concatenated payload sequence of a +// capture contains want (used where multiple datagrams share one writer). +func waitSubstringFrames(t *testing.T, path, proto, want string) { + t.Helper() + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + got, err := transportPayloads(path, proto) + if err == nil && strings.Contains(strings.Join(got, "|"), want) { + return + } + time.Sleep(20 * time.Millisecond) } + got, _ := transportPayloads(path, proto) + t.Fatalf("payload sequence never contained %q (proto %s), last saw %q", want, proto, strings.Join(got, "|")) +} + +// transportPayloads extracts transport-layer payloads from a pcapng file, +// or an error if the file is empty or unreadable. +func transportPayloads(path, proto string) ([]string, error) { + f, err := os.Open(path) + if err != nil { + return nil, err + } + defer f.Close() + r, err := pcapgo.NewNgReader(f, pcapgo.DefaultNgReaderOptions) + if err != nil { + return nil, err + } + var out []string + for { + data, _, err := r.ReadPacketData() + if err != nil { + break + } + pkt := gopacket.NewPacket(data, layers.LinkTypeEthernet, gopacket.Default) + var payload []byte + if proto == "udp" { + if u, ok := pkt.Layer(layers.LayerTypeUDP).(*layers.UDP); ok { + payload = u.Payload + } + } else if t, ok := pkt.Layer(layers.LayerTypeTCP).(*layers.TCP); ok { + payload = t.Payload + } + out = append(out, string(payload)) + } + return out, nil } func TestTCPService(t *testing.T) { - t.Run("echo and capture", func(t *testing.T) { - dir := t.TempDir() + t.Run("echo over tcp with pcapng capture", func(t *testing.T) { conf := echoConfig(t) - conf.CaptureDir = dir + run, path := testRun(t) - svc, err := NewService(conf, nil, testLogger()) + svc, err := NewService(conf, nil, testLogger(), run) if err != nil { t.Fatalf("NewService: %v", err) } @@ -163,8 +201,9 @@ func TestTCPService(t *testing.T) { } _ = conn.Close() - if got := captureFile(t, dir, conf.Name); got != "hello\n" { - t.Fatalf("capture content %q", got) + // produce a pcapng file with the exchanged payload + if _, err := transportPayloads(path, "tcp"); err != nil { + t.Fatalf("capture: %v", err) } }) @@ -172,7 +211,7 @@ func TestTCPService(t *testing.T) { conf := echoConfig(t) conf.ReadTimeout = 200 * time.Millisecond - svc, err := NewService(conf, nil, testLogger()) + svc, err := NewService(conf, nil, testLogger(), nil) if err != nil { t.Fatalf("NewService: %v", err) } @@ -189,8 +228,15 @@ func TestTCPService(t *testing.T) { }) t.Run("script errors don't kill the listener", func(t *testing.T) { - conf := luaConfig(t, "isotest", "isolated.lua") - svc, err := NewService(conf, nil, testLogger()) + conf := Config{ + Name: "isotest", + Network: "tcp", + Addr: freePort(t, "tcp"), + HandlerSpec: "lua:isolated.lua", + BaseDir: "../handler/testdata", + ReadTimeout: 5 * time.Second, + } + svc, err := NewService(conf, nil, testLogger(), nil) if err != nil { t.Fatalf("NewService: %v", err) } @@ -216,54 +262,59 @@ func TestTCPService(t *testing.T) { }) } -func TestTCPServiceTLS(t *testing.T) { - t.Run("echo over TLS", func(t *testing.T) { +func TestTCPCapture(t *testing.T) { + t.Run("tcp listener produces pcapng with frames", func(t *testing.T) { conf := echoConfig(t) - conf.TLS = &tlsprovider.Config{} + run, path := testRun(t) - svc, err := NewService(conf, nil, testLogger()) + svc, err := NewService(conf, nil, testLogger(), run) if err != nil { t.Fatalf("NewService: %v", err) } startService(t, svc) - conn := dialTLS(t, conf.Addr, "localhost") - defer func() { _ = conn.Close() }() - if _, err := conn.Write([]byte("secure")); err != nil { + conn := dialTCP(t, conf.Addr) + if _, err := conn.Write([]byte("hello\n")); err != nil { t.Fatalf("Write: %v", err) } buf := make([]byte, 6) if _, err := io.ReadFull(conn, buf); err != nil { t.Fatalf("ReadFull: %v", err) } - if string(buf) != "secure" { - t.Fatalf("expected echo, got %q", buf) - } + _ = conn.Close() + + waitTransportFrames(t, path, "tcp", + []string{"", "", "hello\n", "hello\n", "", ""}) }) +} - t.Run("SNI visible to script", func(t *testing.T) { - conf := luaConfig(t, "snitest", "sni.lua") +func TestTCPServiceTLS(t *testing.T) { + t.Run("echo over TLS", func(t *testing.T) { + conf := echoConfig(t) conf.TLS = &tlsprovider.Config{} - svc, err := NewService(conf, nil, testLogger()) + svc, err := NewService(conf, nil, testLogger(), nil) if err != nil { t.Fatalf("NewService: %v", err) } startService(t, svc) - conn := dialTLS(t, conf.Addr, "c2.evil.example") + conn := dialTLS(t, conf.Addr, "localhost") defer func() { _ = conn.Close() }() - buf := make([]byte, len("sni:c2.evil.example")) + if _, err := conn.Write([]byte("secure")); err != nil { + t.Fatalf("Write: %v", err) + } + buf := make([]byte, 6) if _, err := io.ReadFull(conn, buf); err != nil { t.Fatalf("ReadFull: %v", err) } - if string(buf) != "sni:c2.evil.example" { - t.Fatalf("expected SNI reply, got %q", buf) + if string(buf) != "secure" { + t.Fatalf("expected echo, got %q", buf) } }) } -func TestUDPService(t *testing.T) { +func TestUDPCapture(t *testing.T) { exchange := func(t *testing.T, addr, payload, want string) { t.Helper() server, err := net.ResolveUDPAddr("udp", addr) @@ -294,23 +345,27 @@ func TestUDPService(t *testing.T) { t.Fatalf("no reply for %q", payload) } - t.Run("echo", func(t *testing.T) { + t.Run("udp echo produces pcapng", func(t *testing.T) { conf := Config{ Name: "udpecho", Network: "udp", Addr: freePort(t, "udp"), HandlerSpec: "builtin:echo", - ReadTimeout: 5 * time.Second, + ReadTimeout: 150 * time.Millisecond, + Capture: true, } - svc, err := NewService(conf, nil, testLogger()) + run, path := testRun(t) + svc, err := NewService(conf, nil, testLogger(), run) if err != nil { t.Fatalf("NewService: %v", err) } startService(t, svc) - exchange(t, conf.Addr, "query", "query") + exchange(t, conf.Addr, "ping", "ping") + + waitSubstringFrames(t, path, "udp", "ping|ping") }) - t.Run("lua packets", func(t *testing.T) { + t.Run("udp lua packets", func(t *testing.T) { conf := Config{ Name: "udplua", Network: "udp", @@ -319,7 +374,7 @@ func TestUDPService(t *testing.T) { BaseDir: "../handler/testdata", ReadTimeout: 5 * time.Second, } - svc, err := NewService(conf, nil, testLogger()) + svc, err := NewService(conf, nil, testLogger(), nil) if err != nil { t.Fatalf("NewService: %v", err) } @@ -328,37 +383,6 @@ func TestUDPService(t *testing.T) { }) } -// TestCaptureStoreEviction verifies idle UDP writers are swept so capture -// files don't accumulate open handles for the life of the listener -func TestCaptureStoreEviction(t *testing.T) { - cs, err := capture.NewStore(t.TempDir(), "evict") - if err != nil { - t.Fatalf("NewStore: %v", err) - } - store := &captureStore{store: cs, idle: 30 * time.Millisecond} - t.Cleanup(func() { store.closeAll() }) - - _, err = store.writer("10.0.0.1:1") - if err != nil { - t.Fatalf("writer a: %v", err) - } - time.Sleep(60 * time.Millisecond) - - if _, err := store.writer("10.0.0.2:2"); err != nil { // sweeps the idle writer - t.Fatalf("writer b: %v", err) - } - if len(store.entries) != 1 { - t.Fatalf("expected idle writer to be evicted, %d entries remain", len(store.entries)) - } - - if _, err := store.writer("10.0.0.1:1"); err != nil { - t.Fatalf("writer a again: %v", err) - } - if len(store.entries) != 2 { - t.Fatalf("expected 2 entries, got %d", len(store.entries)) - } -} - func TestStartWithCancelledContext(t *testing.T) { for _, network := range []string{"tcp", "udp"} { conf := Config{ @@ -368,7 +392,7 @@ func TestStartWithCancelledContext(t *testing.T) { HandlerSpec: "builtin:sink", ReadTimeout: 5 * time.Second, } - svc, err := NewService(conf, nil, testLogger()) + svc, err := NewService(conf, nil, testLogger(), nil) if err != nil { t.Fatalf("NewService: %v", err) } @@ -420,7 +444,7 @@ func TestNewServiceValidation(t *testing.T) { t.Run(tc.name, func(t *testing.T) { conf := base() tc.mutate(&conf) - _, err := NewService(conf, nil, testLogger()) + _, err := NewService(conf, nil, testLogger(), nil) if err == nil || !strings.Contains(err.Error(), tc.wantErr) { t.Fatalf("expected error containing %q, got: %v", tc.wantErr, err) } diff --git a/internal/listener/service.go b/internal/listener/service.go index abca52f..a721950 100644 --- a/internal/listener/service.go +++ b/internal/listener/service.go @@ -3,8 +3,6 @@ package listener import ( "fmt" "log/slog" - "sync" - "time" "github.com/lachlanharrisdev/gonetsim/internal/capture" "github.com/lachlanharrisdev/gonetsim/internal/handler" @@ -12,7 +10,7 @@ import ( "github.com/lachlanharrisdev/gonetsim/internal/state" ) -func NewService(conf Config, global *state.Store, logger *slog.Logger) (service.Service, error) { +func NewService(conf Config, global *state.Store, logger *slog.Logger, run *capture.Run) (service.Service, error) { if global == nil { global = state.NewStore(nil) } @@ -25,83 +23,11 @@ func NewService(conf Config, global *state.Store, logger *slog.Logger) (service. } log := service.NewPrefixedLogger(logger, conf.Name) - store := &captureStore{} - if conf.Capture { - baseDir := conf.CaptureDir - if baseDir == "" { - baseDir = capture.DefaultDir - } - cs, err := capture.NewStore(baseDir, conf.Name) - if err != nil { - return nil, fmt.Errorf("listener %s: %w", conf.Name, err) - } - store.store = cs + if !conf.Capture { + run = nil } if conf.Network == "udp" { - store.idle = conf.ReadTimeout - return &udpService{conf: conf, handler: h, log: log, store: store, global: global}, nil - } - return &tcpService{conf: conf, handler: h, log: log, store: store, global: global}, nil -} - -type captureStore struct { - store *capture.Store - idle time.Duration - - mu sync.Mutex - entries map[string]*captureEntry -} - -type captureEntry struct { - w *capture.Writer - last time.Time -} - -func (cs *captureStore) writer(key string) (*capture.Writer, error) { - if cs.store == nil { - return nil, nil - } - cs.mu.Lock() - defer cs.mu.Unlock() - - now := time.Now() - if cs.idle > 0 { - for k, e := range cs.entries { - if now.Sub(e.last) > cs.idle { - _ = e.w.Close() - delete(cs.entries, k) - } - } - } - if e, ok := cs.entries[key]; ok { - e.last = now - return e.w, nil - } - w, err := cs.store.Conn(key, now) - if err != nil { - return nil, err - } - if cs.entries == nil { - cs.entries = make(map[string]*captureEntry) - } - cs.entries[key] = &captureEntry{w: w, last: now} - return w, nil -} - -func (cs *captureStore) closeAll() { - cs.mu.Lock() - defer cs.mu.Unlock() - for _, e := range cs.entries { - _ = e.w.Close() - } - cs.entries = nil -} - -func (cs *captureStore) release(key string) { - cs.mu.Lock() - defer cs.mu.Unlock() - if e, ok := cs.entries[key]; ok { - delete(cs.entries, key) - _ = e.w.Close() + return &udpService{conf: conf, handler: h, log: log, run: run, global: global}, nil } + return &tcpService{conf: conf, handler: h, log: log, run: run, global: global, idle: conf.ReadTimeout}, nil } diff --git a/internal/listener/tcp.go b/internal/listener/tcp.go index 67f3b55..a357e48 100644 --- a/internal/listener/tcp.go +++ b/internal/listener/tcp.go @@ -10,52 +10,72 @@ import ( "sync" "time" + "github.com/lachlanharrisdev/gonetsim/internal/capture" "github.com/lachlanharrisdev/gonetsim/internal/handler" + "github.com/lachlanharrisdev/gonetsim/internal/netx" "github.com/lachlanharrisdev/gonetsim/internal/state" ) type tcpService struct { conf Config - handler handler.Handler + handler handler.TCPHandler log *slog.Logger - store *captureStore + run *capture.Run global *state.Store + idle time.Duration + mu sync.Mutex + ln net.Listener conns connSet wg sync.WaitGroup } func (s *tcpService) Name() string { return s.conf.Name } -func (s *tcpService) Stop(_ context.Context) error { return nil } - -func (s *tcpService) Start(ctx context.Context) error { - ln, err := net.Listen("tcp", s.conf.Addr) - if err != nil { - return err +func (s *tcpService) Stop(_ context.Context) error { + s.mu.Lock() + ln := s.ln + s.mu.Unlock() + if ln != nil { + _ = ln.Close() } - defer func() { _ = ln.Close() }() + s.conns.closeAll() + return nil +} +func (s *tcpService) Start(ctx context.Context) error { + var tlsConf *tls.Config if s.conf.TLS != nil { - tlsConf, err := s.conf.TLS.TLSConfig() + var err error + tlsConf, err = s.conf.TLS.TLSConfig() if err != nil { return err } - ln = tls.NewListener(ln, tlsConf) } + iface, err := s.run.NewInterface("gonetsim " + s.conf.Name + " tcp") + if err != nil { + return err + } + ln, err := netx.ListenTCP(s.conf.Addr, s.run, iface, tlsConf) + if err != nil { + return err + } + defer func() { _ = ln.Close() }() - done := make(chan struct{}) - defer close(done) - go func() { - select { - case <-ctx.Done(): - _ = ln.Close() - case <-done: - } + s.mu.Lock() + s.ln = ln + s.mu.Unlock() + defer func() { + s.mu.Lock() + s.ln = nil + s.mu.Unlock() }() + done := netx.CloseOnCancel(ctx, ln) + defer done() + s.log.Info("listening", "on", s.conf.Addr, "handler", s.conf.HandlerSpec) - if err := s.accept(ctx, ln); err != nil && !errors.Is(err, net.ErrClosed) && ctx.Err() == nil { + if err := s.accept(ctx, ln); err != nil && !netx.IsExpectedClose(err, ctx) { return err } @@ -83,25 +103,22 @@ func (s *tcpService) accept(ctx context.Context, ln net.Listener) error { func (s *tcpService) handleConn(ctx context.Context, conn net.Conn) { defer func() { _ = conn.Close() }() - remote := conn.RemoteAddr().String() - defer s.store.release(remote) - - w, err := s.store.writer(remote) - if err != nil { - s.log.Warn("capture unavailable", "remote", remote, "err", err) + var env *capture.Session + if cc, ok := conn.(*capture.Conn); ok { + env = cc.Session() } conn = newIdleConn(conn, s.conf.ReadTimeout) - env := handler.Env{Logger: s.log, Capture: w, IdleTimeout: s.conf.ReadTimeout, Global: s.global} - err = s.handler.HandleTCP(ctx, conn, env) + henv := handler.Env{Logger: s.log, Capture: env, IdleTimeout: s.conf.ReadTimeout, Global: s.global} + err := s.handler.HandleTCP(ctx, conn, henv) switch { case err == nil, errors.Is(err, net.ErrClosed), errors.Is(err, os.ErrDeadlineExceeded), errors.Is(err, context.Canceled): - s.log.Debug("connection closed", "remote", remote) + s.log.Debug("connection closed", "remote", conn.RemoteAddr().String()) default: - s.log.Info("connection handler error", "remote", remote, "err", err) + s.log.Info("connection handler error", "remote", conn.RemoteAddr().String(), "err", err) } } diff --git a/internal/listener/udp.go b/internal/listener/udp.go index 10ea352..26cbc45 100644 --- a/internal/listener/udp.go +++ b/internal/listener/udp.go @@ -4,10 +4,12 @@ import ( "context" "errors" "log/slog" - "net" "os" + "sync" + "github.com/lachlanharrisdev/gonetsim/internal/capture" "github.com/lachlanharrisdev/gonetsim/internal/handler" + "github.com/lachlanharrisdev/gonetsim/internal/netx" "github.com/lachlanharrisdev/gonetsim/internal/state" ) @@ -17,42 +19,60 @@ const maxPacketSize = 65535 // receive order and scripts never run concurrently. type udpService struct { conf Config - handler handler.Handler + handler handler.UDPHandler log *slog.Logger - store *captureStore + run *capture.Run global *state.Store + + mu sync.Mutex + pc *capture.PacketConn } func (s *udpService) Name() string { return s.conf.Name } -func (s *udpService) Stop(_ context.Context) error { return nil } +func (s *udpService) Stop(_ context.Context) error { + s.mu.Lock() + pc := s.pc + s.mu.Unlock() + if pc != nil { + _ = pc.Close() + pc.CloseAll() + } + return nil +} func (s *udpService) Start(ctx context.Context) error { - pc, err := net.ListenPacket("udp", s.conf.Addr) + iface, err := s.run.NewInterface("gonetsim " + s.conf.Name + " udp") if err != nil { return err } - defer func() { _ = pc.Close() }() + rec, err := netx.ListenUDP(s.conf.Addr, s.run, iface, s.conf.ReadTimeout) + if err != nil { + return err + } + defer func() { _ = rec.Close() }() - done := make(chan struct{}) - defer close(done) - go func() { - select { - case <-ctx.Done(): - _ = pc.Close() - case <-done: - } + s.mu.Lock() + s.pc = rec + s.mu.Unlock() + defer func() { + s.mu.Lock() + s.pc = nil + s.mu.Unlock() }() + done := netx.CloseOnCancel(ctx, rec) + defer done() + s.log.Info("listening", "on", s.conf.Addr, "handler", s.conf.HandlerSpec, "net", "udp") - if err := s.readLoop(ctx, pc); err != nil && !errors.Is(err, net.ErrClosed) && ctx.Err() == nil { + if err := s.readLoop(ctx, rec); err != nil && !netx.IsExpectedClose(err, ctx) { return err } - s.store.closeAll() + rec.CloseAll() return nil } -func (s *udpService) readLoop(ctx context.Context, pc net.PacketConn) error { +func (s *udpService) readLoop(ctx context.Context, pc *capture.PacketConn) error { buf := make([]byte, maxPacketSize) for { n, remote, err := pc.ReadFrom(buf) @@ -63,11 +83,7 @@ func (s *udpService) readLoop(ctx context.Context, pc net.PacketConn) error { data := make([]byte, n) copy(data, buf[:n]) - w, err := s.store.writer(remote.String()) - if err != nil { - s.log.Warn("capture unavailable", "remote", remote.String(), "err", err) - } - env := handler.Env{Logger: s.log, Capture: w, Global: s.global} + env := handler.Env{Logger: s.log, Capture: pc.SessionFor(remote), Global: s.global} reply, err := s.handler.HandleUDP(ctx, data, remote, env) if err != nil { diff --git a/internal/netx/netx.go b/internal/netx/netx.go new file mode 100644 index 0000000..cfc0051 --- /dev/null +++ b/internal/netx/netx.go @@ -0,0 +1,107 @@ +package netx + +import ( + "context" + "crypto/tls" + "errors" + "fmt" + "io" + "net" + "strconv" + "strings" + "time" + + "github.com/lachlanharrisdev/gonetsim/internal/capture" +) + +func ParseAddr(addr string) (string, error) { + if addr == "" { + return "", fmt.Errorf("listen address is required") + } + if _, err := net.ResolveTCPAddr("tcp", addr); err != nil { + return "", fmt.Errorf("invalid listen address %q (expected host:port): %w", addr, err) + } + return addr, nil +} + +func ValidateNetwork(network string, allowed ...string) error { + n := strings.ToLower(strings.TrimSpace(network)) + for _, a := range allowed { + if n == a { + return nil + } + } + return fmt.Errorf("network must be one of: %s", strings.Join(allowed, ", ")) +} + +func ValidateStatus(code int) error { + if code != 0 && (code < 100 || code > 599) { + return fmt.Errorf("status code must be 0 or between 100 and 599, was %d", code) + } + return nil +} + +func DisplayNetwork(network string) string { + switch strings.ToLower(strings.TrimSpace(network)) { + case "both": + return "udp+tcp" + case "tcp": + return "tcp" + default: + return "udp" + } +} + +func ParsePort(addr string) (int, bool) { + _, portStr, err := net.SplitHostPort(addr) + if err != nil { + return 0, false + } + port, err := strconv.Atoi(portStr) + if err != nil { + return 0, false + } + return port, true +} + +func ListenTCP(addr string, run *capture.Run, iface int, tlsCfg *tls.Config) (net.Listener, error) { + ln, err := net.Listen("tcp", addr) + if err != nil { + return nil, err + } + ln = capture.NewConnListener(ln, run, iface) + if tlsCfg != nil { + ln = tls.NewListener(ln, tlsCfg) + } + return ln, nil +} + +func ListenUDP(addr string, run *capture.Run, iface int, idle time.Duration) (*capture.PacketConn, error) { + pc, err := net.ListenPacket("udp", addr) + if err != nil { + return nil, err + } + return capture.NewPacketConn(pc, run, iface, idle), nil +} + +func CloseOnCancel(ctx context.Context, c io.Closer) (stop func()) { + stopped := make(chan struct{}) + go func() { + select { + case <-ctx.Done(): + _ = c.Close() + case <-stopped: + } + }() + return func() { close(stopped) } +} + +func IsExpectedClose(err error, ctx context.Context) bool { + if err == nil { + return true + } + if errors.Is(err, net.ErrClosed) || errors.Is(err, context.Canceled) { + return true + } + return ctx.Err() != nil +} diff --git a/internal/observability/logging.go b/internal/observability/logging.go index 13a2538..30693e0 100644 --- a/internal/observability/logging.go +++ b/internal/observability/logging.go @@ -9,13 +9,16 @@ import ( "github.com/lmittmann/tint" "github.com/mattn/go-colorable" "github.com/mattn/go-isatty" - - "github.com/lachlanharrisdev/gonetsim/internal/config" ) -func NewLogger(cfg config.LoggingConfig) (*slog.Logger, error) { +type Options struct { + Format string + Level string +} + +func NewLogger(cfg Options) (*slog.Logger, error) { level := parseLevel(cfg.Level) - if strings.ToLower(strings.TrimSpace(cfg.LogFormat)) == "json" { + if strings.ToLower(strings.TrimSpace(cfg.Format)) == "json" { return slog.New(slog.NewJSONHandler(os.Stdout, &slog.HandlerOptions{Level: level})), nil } diff --git a/internal/service/manager_test.go b/internal/service/manager_test.go index 600821f..0a7cc3e 100644 --- a/internal/service/manager_test.go +++ b/internal/service/manager_test.go @@ -1,13 +1,22 @@ +////---------------------------------------------------------------------------- +// NOTICE: to save development time, test files (including this) have been +// generated with LLMs. The author(s) do not claim credit for these tests +// and exist purely for maximising code quality and reliability +// +// For more information please see `/.github/AI_USAGE.md` +//----------------------------------------------------------------------------// + package service import ( "context" "errors" - "io" "log/slog" "strings" "testing" "time" + + "github.com/lachlanharrisdev/gonetsim/internal/testutil" ) type fakeService struct { @@ -34,7 +43,7 @@ func (f *fakeService) Stop(ctx context.Context) error { } func discardLogger() *slog.Logger { - return slog.New(slog.NewTextHandler(io.Discard, nil)) + return testutil.Logger() } func TestRunServices_PropagatesStartError(t *testing.T) { diff --git a/internal/state/state.go b/internal/state/state.go index fa46b4e..fefb5dd 100644 --- a/internal/state/state.go +++ b/internal/state/state.go @@ -3,8 +3,6 @@ package state import ( "errors" "fmt" - "strconv" - "strings" "sync" ) @@ -89,38 +87,3 @@ func (s *Store) Delete(key string) { delete(s.data, key) } } - -// could move to a shared utils package but not necessary yet -func ParseSize(s string) (int64, error) { - s = strings.TrimSpace(strings.ToLower(s)) - if n, err := strconv.ParseInt(s, 10, 64); err == nil { - if n <= 0 { - return 0, fmt.Errorf("size must be positive") - } - return n, nil - } - - var mult int64 - switch { - case strings.HasSuffix(s, "kib"): - mult, s = 1<<10, s[:len(s)-3] - case strings.HasSuffix(s, "mib"): - mult, s = 1<<20, s[:len(s)-3] - case strings.HasSuffix(s, "gib"): - mult, s = 1<<30, s[:len(s)-3] - case strings.HasSuffix(s, "k"): - mult, s = 1<<10, s[:len(s)-1] - case strings.HasSuffix(s, "m"): - mult, s = 1<<20, s[:len(s)-1] - case strings.HasSuffix(s, "g"): - mult, s = 1<<30, s[:len(s)-1] - default: - return 0, fmt.Errorf("invalid size %q (expected e.g. 64MiB)", s) - } - - n, err := strconv.ParseInt(strings.TrimSpace(s), 10, 64) - if err != nil || n <= 0 { - return 0, fmt.Errorf("invalid size %q (expected e.g. 64MiB)", s) - } - return n * mult, nil -} diff --git a/internal/state/state_test.go b/internal/state/state_test.go index da50b4f..ac47d83 100644 --- a/internal/state/state_test.go +++ b/internal/state/state_test.go @@ -1,3 +1,11 @@ +////---------------------------------------------------------------------------- +// NOTICE: to save development time, test files (including this) have been +// generated with LLMs. The author(s) do not claim credit for these tests +// and exist purely for maximising code quality and reliability +// +// For more information please see `/.github/AI_USAGE.md` +//----------------------------------------------------------------------------// + package state import ( @@ -76,32 +84,4 @@ func TestState(t *testing.T) { t.Fatalf("empty value should be allowed: %v", err) } }) - - t.Run("parse size", func(t *testing.T) { - cases := []struct { - in string - want int64 - wantErr bool - }{ - {"64MiB", 64 << 20, false}, - {"64mib", 64 << 20, false}, - {"512K", 512 << 10, false}, - {"1GiB", 1 << 30, false}, - {"4096", 4096, false}, - {"", 0, true}, - {"64GiB", 64 << 30, false}, - {"abc", 0, true}, - {"-1MiB", 0, true}, - {"64TiB", 0, true}, - } - for _, tc := range cases { - got, err := ParseSize(tc.in) - if tc.wantErr && err == nil { - t.Errorf("ParseSize(%q): expected error", tc.in) - } - if !tc.wantErr && (err != nil || got != tc.want) { - t.Errorf("ParseSize(%q) = %d, %v; want %d", tc.in, got, err, tc.want) - } - } - }) } diff --git a/internal/testutil/testutil.go b/internal/testutil/testutil.go new file mode 100644 index 0000000..8ab96aa --- /dev/null +++ b/internal/testutil/testutil.go @@ -0,0 +1,35 @@ +package testutil + +import ( + "io" + "log/slog" + "net" + "testing" +) + +func Logger() *slog.Logger { + return slog.New(slog.NewTextHandler(io.Discard, nil)) +} + +func FreeTCPAddr(t *testing.T) string { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("Listen: %v", err) + } + defer func() { _ = ln.Close() }() + return ln.Addr().String() +} + +func FreePort(t *testing.T, network string) string { + t.Helper() + if network == "udp" { + pc, err := net.ListenPacket("udp", "127.0.0.1:0") + if err != nil { + t.Fatalf("ListenPacket: %v", err) + } + defer func() { _ = pc.Close() }() + return pc.LocalAddr().String() + } + return FreeTCPAddr(t) +} diff --git a/internal/tlsprovider/config.go b/internal/tlsprovider/config.go index 70d8def..1c8c2d9 100644 --- a/internal/tlsprovider/config.go +++ b/internal/tlsprovider/config.go @@ -30,12 +30,16 @@ type Config struct { } func (c Config) Validate() error { - if (c.CertFile == "") != (c.KeyFile == "") { // temu xor + if (c.CertFile == "") != (c.KeyFile == "") { return errors.New("cert and key must be set together") } return nil } +func DefaultPaths(configDir string) (cert, key string) { + return filepath.Join(configDir, PersistedCertFileName), filepath.Join(configDir, PersistedKeyFileName) +} + func (c Config) TLSConfig() (*tls.Config, error) { if err := c.Validate(); err != nil { return nil, err @@ -127,7 +131,6 @@ func (c Config) regeneratePersistedPair() error { return nil } -// certExpired reports whether the leaf certificate has passed its NotAfter time func certExpired(cert tls.Certificate) bool { if len(cert.Certificate) == 0 { return false diff --git a/internal/tlsprovider/tls_test.go b/internal/tlsprovider/tls_test.go index a70dbd2..4aae8f2 100644 --- a/internal/tlsprovider/tls_test.go +++ b/internal/tlsprovider/tls_test.go @@ -1,3 +1,11 @@ +////---------------------------------------------------------------------------- +// NOTICE: to save development time, test files (including this) have been +// generated with LLMs. The author(s) do not claim credit for these tests +// and exist purely for maximising code quality and reliability +// +// For more information please see `/.github/AI_USAGE.md` +//----------------------------------------------------------------------------// + package tlsprovider import ( diff --git a/smtp.pcap b/smtp.pcap deleted file mode 100644 index 931b43b3b878dc559b7423df26363bfe75c6a512..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 27850 zcmeHw3y@padEO-{>aq6Pq^EIdM;)Do#Arz^u)w}Zf{zdjU?0>zpj|8}$Fsy#Grk7^ ze09Qq&xUc98NAQ>?m^{uN*%g#?yCn4BRRLQtY%Xh z+{2%^-Z%wFeDzRrc<9@tWa8;gN@grj^0TP}?|l9X5Aw0!`SPLXiQ1D-TyOk)eCDeM z*3&~T3aneVg@sj#yc4T3x$~;bS;e9Qyw!UmV`HJn6kZeKp~&R1dqI3Ww>p1!QcaFc zj2%B74(qQ{^Qq**{L=i=to|HZOV6p)LTr9XEiBK@FL@89V=L+O!m4~ybIApdG#~yK z*BjHA312;wo;ZAqQ2GzLDE-+_zmj=1b?Eiq|K);V$avmJB+$#>88qeqjISO7MTcGz z0)4EDK(B(>e@M2?LX(Y;PpI>EE)}z-$wKDB+LT_zTVEbD(g4X<52iDRj|is!ybIGm z{PZhX1DJj$@EeA4_xSm(?Ci`|Zg!?|$J}%wH+Qa*8DF`)P%7332fhg03*Xq+5BKi~ z?jP;MJ!=5>7m54XNsas3?9Aotcy1$Gx&ycq?kB(52De*HM-PVt_uuKl{YBvZ4&eT= z^wWlsojYGQOEdNziR9K|?^t(Bm!N)M0 z$*@|!T|=9G>}|%qK$))|yq6Cj7rg&rC*J9I1MlO&yYb}}!&r%@(ki_StgGhz7ju?V zgjO80%V)RDg5yAoYSjM(3;gwOUSmMv%BzRzr-oh@)Vr8IdmO08i24g}*Qm!p_zLvn z-~eGh_rf(Tt<}vmp~I(T0#9{KAchIdU;_4;7lBPI9n-Krb7;VTWtUeEe)kQ1SMaNL zk#+|7%@V)U889}?K(bu3s;XqUuDNcFD65>a%j>FP*9s~S3MlZ7sfu>lQkgAf*IcWZ z*P~zeYXio2_Z#x+UW)x+KfKJDeE%=HW-^PJJdir@dgLn~FpQb1Q;Mpgf)3AAHF4Ps!GRkxPQ9Ent6Qf7LHs#=amm`OrEteZjL6r-#c(7Zo0yyq zhr%Z&r%sH|gad;EcLo<>L9HrgVP~t3>*Q-{IiI()mRhdYijH$hMaDROx=^*uTyVae zvo@nDmaWXJTcs1z8~5LTcQRqEL~bt}w{I(~ZJl?*F=`qg-vzV0UdPHuvGT#XJ2u#%`8nv)J*Eqikd0k zc0tX}#8YaneAhX3Zuyw2Zm*Z7)a_ff3u<8|JFgb*U#p#roXoCI7hS$FWhbza|y#V2BSc|YYk=5cb>aSLYT;zs1;v6JtM z$718hW+qPrBgbJ@Cnn>^g2yK(!ogTL)V{?nKh4tx5Y<41cve%UY7vWz-%`T=-t~I8YVl#Kd)?zV` zSE#7zXXm2ojiEL;z;%SHg@0y?hkW5wc3F}0RNkqU%o^OP!2$Bp8|0?ens@W4S%eBG z!%%EkSK-tMIdyVdPTc{Uou25%rUr2strShWOcb&Ov+7#4(}B7hG~KLiLlk_DP1F0# zTh(B)oOP(IqUu=2u5o!-fve{{$&@-@NK8@%CEb8SPHICP;_t3*dfy%~m@Gh~=Es-GK$;hq?cjEy>__8GT5 z-DosIjq#9ET^~)ajIJc(K~ixd9Dr{^l0sqP-KoW7I;Q5*=~OUz$J+dP0%-CYw0S_q zmzUDXr8F-Ho65S9#6FvjCz^apfF-MDs?A9E3e**Gor;<0-ecRkZkKC>R7&j~-Ps0HP9B-Y(np>s)no6c^6s@pUk7X2B~r zywCekvWr{MK-?^4sX%~~NqW~`WH$F+iy=$M*NlWH8V2{nbE_>_B5MqrHFzRR0F6$JIbNWg{i2)b6)&IdSKGYa44(ui_F77jg! z1+Vfe8Qpb)vILYx1jV(j=!p{|&`!~*Ml;2_=u7fMFlS{QtTP<~*m;3r=>axuhUtip zd&Zb?5EwS-8e*pX$Ev_dDqejw86wG`TZkVy)(N^s=L8dUow8Y^2v)!fO#=2O79eG1 z+6|~qr@={cej!pVUZ{X4O*lQorS(+Umr;XHG#!;vRL%iNcMj zAU={wxN0Mc=}}T4jI{t(&d!z0O@0vxk4**Cs6gh^r-t4`mpOkgMx&TsWrU)+<!AX0%4R13}Zc7#6dL4peAqyMjHWVX9m?)!yu#;u%u19$J_2E5VHws95S=%YUyqVPMm13ZfL(6kuk zT6%A3gXHbc;H1!qZ0ptR6o9ljT@u!GI%#6z0!zsYtD)Lvt-puX>(lrqHen-`2n{Kr zRLdI^tpsnaXOO`GFYrQ&^924vg&c)-D{mp?xDE$z(?Sq)RLz$vRc8bCloISzIq}Zt z5HrC<)hm+6ghve-sab9faaMV~0b^aNFz)iRo|(E`%n9J>)+n7IkB3eWmU0VDy_iF)x1p+*%-t@#HM5A>V+9n5 z29LAksyu>ol`^xJKw5uVoQGRbZnkP!NDBaGCql)kIGv>TAX}h*?-F$r52k9n zNzKON>Wmr>jq!5A;u>$5aA<%ff~2bD%>1%_y}h zn8&G*sBZN17`RQTAY7s64E=>h- z7ycio=D*Ni^RI~J|9+?Dx3a$G|JK8XF@jL5P~>ie{u!)g>?*t>@6$x{-DX8*x-$7O}G%k%90UcBkE(#@-~^y%;C`0LIS^H?B$o9$lPxk38Q!0)RxzrV zR&y0{K`wdI_S6z@>7z*|d0k65r5Xy-wKH}R9w_%;N(g)Nc!w(+CapNOXw0Tq_+o&W zh!DVmURT3ad%XZ2$7anXozrVj9$F7knvvY2AOsdQe1RR1*(bDQC9I8L{I=WFj9rDM zOxq=>tL1XuMgmyt0PP!&ko;9wWzDjPXBL@NY$O4pj8$t;chrR#AP{9xd<~VUIt@If z*OU5_D!6EgR?1sutzNZUq&%*!NGAHyPV>PQhQSJEi**c?vnl^hbxTB!R2BUVblXtK zLj0)HRR))lYJ*1))`dpW6v+!kph&B=dW8ry?T^6f2E;pBI=NvNi$W+ePaPJ+X_SkO zN&OD^Df1FCLS=ubd#IG*>bj(76Lm-r7SPA0rC|7A&f3UUDj|FsnvTV9UrTZHX?vYA zfgvzpn>3J95xK&MxNsER8kwh2VA_O+6rbeboo(PauP-;Tg+V5()W|BZ4{fZHv(2Kj z4!kwByD3DT#sp!QgwH`)W>0gTit?9^`J5J?VQJQNxNYTxRho;?qBUfv&Xmg4`A94d zt+_TO9uEqSBo-FjD3+h8Bl3i)s`WB#`8w?<6%OI^QG8@R+A{=EwFWWB)?HY0xGa|t z#uu%SOiAC~L=&1r@RV7^7$u5GajZ55_Vae#q?Ife)-+k(u&Yj)?gVM+X(^9{!&qAo zmof^v!SkqAZTJ>U^K0o01hdrqw?8&uEbTSq)wRp;%-*xfcVRy#s zy!`!NN*a0(;DKv@u;U)UP2btw19&s<0g$^s0PcDLyGVTX>>_>X@T&L~KkD==(#>6@ z)#xq^cNb}!4K9YdySqq|rrq5|`hRN|i5h?QV}1RJcZ)wU(()(vf3E3IU}p*b z#Ev^lU-_loouxPP&JwxncR${jyMA7{>$Vnm-E1|v3)yRM*N)liEfj3+X0LZn%K!hK zy(V`>AMeXuHwt$hZE@EZ&NaCU88UF!jv2C}ukL2Z-pm=YKLvMv|HFG)wSnD*pTk4n zlX|{Sb=C7d@W8Woy_Pxoq@L`{L_{ z(N$HlTdSdVYc*QM86(J=?$&B-Ul+1ltFc?FAr&J3$7(gG@qeMe#$OYS|HDp=fAQnK z#{cQdhS6Q|vFq)=dA;4v`qSN7jon%eX?BpBo8zI~T8;nPwHnm;UpsnTzj)yli5Gse zH(t2-@zjCfeV_Uvs zeeBkK?ACoqE#7Y32P+qvRX)1JX1DGG^(p_G>prOQr*7-3@h^zR@2cIn|0ll2|EnJv z#?`7v0$kh@>yl6oi6S=nZ@X(krMCA1%PytliUl^NLh;Rr z)CBouSf~wRX)HQCP}0is7yxE{D=WEB_TtsxphE*)GpKh${e1-`_|4L(-eILQI9@=_ zJZ|@klZh*tl~uuheP?q@QlORcihaB~8<}OZ0;ymXb^BT0`l(~OhIqwtoq9Eknk!Zj zw%56`Vi!oExQDZA-iK~YOR9pq3!ytHdMtEYS3aTYtx<4L)TMiA{8C_4Xw|AtQRPv6 zCY4_-NJWfQn?x|A7+>0343P?FxK%JNtUjr`p;#eNwguju;AuG&k#@(e3r zs21pVvSc)QaU*oZS~bDym{< zrTP)9xywgTkAqd*AdYQT|&xxVB?RWNb!Vj?-cHhb|RQBAUh&M%YG z5)BfmQ{Gn>S^t0LBKog3K<7|FcW{zV(J4yYFbYc7bdk4$TjUHA@L{#e3d*3apwwlo zEcDd{@0xuA$xUglEpn3UD6Mn@l-@L&1KiQ}0zJQmc$jtHb+q?MfoEVD$jAcdLR=4M z>stcz$)}$paaaM5vOk=F;ju9_{)LD8#tYwR5#{uAs_8OGna^Wawr z%L1%E@xq;AJ-T-4D^G~^IP`iXc)&1JK)BojJ=iIcF{MjjhVAm0kR{z|w%3)=yCmQ& z^fp`qUrF(PQQ9jsaaf2iD@p>OtG!m;LJK>zj4z=oCJX@~2*U74L1)gXvua$rR7971 zHM7|~jT5f0xt^tw8@i2xbX^7u@tm7x(@D3kl~Pz?J2NMTD)fjDE3vEAd+X>s8NnOr z{|E#@5v}`l69za>w-a#CcvC2aouEco5s!jlk)`dm=V}4Qum8V8=M13h2Fd_>6skcoDiIFMdnMR@#VRLCsRlAEqwTz#s?Mvz zjrt0uSB439)uLslp=7nXuY3G<3{BQ}@AnWTzbcpVKI~dBkXT=@QRe2U$%$hV6QZ#@ zNAQgTd&1+FXEw-5vqp#(ZPX~PR208m6 zZAB1msF(y>CDt(l%H*NgE7EiU60>dV7>}$#2{w<(KsT!6!YMiqW27 z78RatEdxjHBUeH?T@BduK+7M_cp$59eqfjpO< z&dM;+aUnm#3=||Zx(SXRK=^p*0@8jFOs!R>|AFIjGy;Jm7)!br(CbTh1sWf`uQtU| zd$95It5_j%xkq?i-Fhoxs!QaaMuGi3;O_$L4)`4tM}sn$J`@u9({_OUiM~J0wfE5D znzi2`e9?9dK?!>NFlk5(8Nip|7O*^fKxkkI3a24R*h3;fpLQzc9u=;E>2@HD*{jA= zo*a2@8KI;IE$&{ROT$9~ZQ=u>fpls*a`$+~6{s(wZe?wm2R!XBFxx#ZqJ4e*w#MKO- zrNxF~qMqe@gh|aPSR~J%JcdLMkM4GAxdqzpqQ#H>Gon1r%#mgr71}g_7yNBkyj)Rt zV2x(5rk~bLQW~WIey9Z;+5pAY-niyt;mJ~8eKVy)(srJ4Ewc)*9auJ_THUTl3bEys zol$VE*>^YWyQ4rIO`2@r&!&WCGAAD<$J)*}d_Hck|8=+}}% zqXC-?$4x5-t_U6=lW_8$Ge^e*kWA?>#5GBcx3?cp=|H@_dlzX91qp|N)3Q@=r1@$@ z`={gvY2+hk>Gm$_d^REx1+7oK_AgNobe1|B>$FdaXj=S1yA1D@=Taa}C}m)r#M{$E z0b11fH~+i0_S;{0RpNzD_QnhM|0H!Fao~gBWxVkFuN?^xRsyj4#0z(e_2{aOdtdX+ zm>XWh37<5qjEzX=TcgoiJ;Xs}>KJ+iaQSc@%Sh890fg=!6dNYi=O}d=2zRvE{fXdi zy-r;SRg37mbt{A>d&#Qd+x0?`&69X6q{N*k23hzIVo$&sCPH<&nRzL&J{H7(8q1{> z2)fj802o3+c&56?xCz)#e<0P8Yq|=U1+=GlbM%9{lT=fB56;(dejj)XeQjIB`n!sH zpqVLQl+m`>zW1cL&JR35$AF*Ra8<(Zkd)Xp(C%VkuAUYVn6%1mlt*(RGu@*4cJyeb zbRs6HO}v-npnE2V+~Cle%b*L0fdz(GgWi6dqx2RA(R$aKcFx({LJy!mwuOdGFD>kTVAuS-}-^%jU_s1w>84&TOgjezD#I^;T;Llv^I;xO%GZGv)U8F^L- zGI&}Tr;d(P*qwqG_H|3%WJs2R2zbH*DtWx?NdnRh47v5m6QWk!5L-WS1{6#O;x%be zC&@rQY8^0`;`cl!0=y1@yy!hIxbd?jJn{bqxQqRoK$ONd#B&od2XPkxu%Ufbg3Ovo zk&5Bs`50a@c^GY@I0RxH1{}V&J`Mt|1sO$B9B@C6pwAwT5JV90*6j^a1sD({0BRNN zO`fdKoDlp~3DBwWFZS2??-7mvO6QL9Lr?ngLV3+Fs(KrwAEoJDbSRu7UXN|(x15IT zE`c`{8n25uc?y#U|Fe0R=PF=+PPf?zLD4`OgeRbNxCDBOiU$^r5E8!v*E6jLClVez zdU|OM02?PA(7yvzuo#X%s^c63#tgW91AVqQV#s2!3I9eGstaFQc>)(|<1j`#BN_o6 z(7~t8SRBH`ocCgm6k6dWeAL42v6A3@gl-vQ3M6ZR8_p&HTh8aThhc3(j}?(IXbuZ_ zBoNUHc&aN@2#mmYLzgworH>M{PAoM)#E?f13wzVp9Bpj&+zp5hP7hV8v8h#+JEZPH z*aL8ylseJvg*f^=6!_80+<5?_dJYX}Cn~2sd)v&IVZNPS561LEBw74<99qrnG!_O- z4^gL9$T&DAL;O9`Q+7HS9CKp{qYcyf!(krzJmi9F2$-M}4U|yF1oAk_0{X%WTRhQ# zAi>*$Nf0Y8Xfa`cqo=?hIXhRrl@W)YIzX~)6p1*Qy<+I1S6@_yY+Z2}NkwppI%Wxr zls$F>bihOcx;T#<$dcQ1^op~^We))u=_kE}2yp65noerNw1)h9 zRV}4Qcw~84%=4WoeX2xYXCQmoBEPhfDE+3pF9lT6({1NyiJY& zxkH2f;)NebyztL@k+>N zyVI0h&ot$5$PPV!QWR7Ym5rEzQ|(?p@jhC!Z<%;u+@R+M`BsMeaa@^Ua|37YiCGtO zEvK#MY$XMvDX{+Ye1<&dMVGjxO;Ys=_)adfL0I1k6s`_*Hf+b+76Ik`tdi_30S&;! zxB8rB=r~Cno!o{?#ffSUzKLQ$eAuL%*GE@Y^cEE5TzYXK(2STwk!T79(m2fs;|WKb zNHOT&5g>O+q~srq0Y=pu+P+(fLjgfMtBO7Xhk6hh;|30k0hh|5QSfH~M5%dy{~2~z zM)9AAA7vJ)!g#hF+<2bsJw5p7IJvSQCbIzWNofK&v<}z!vwC%fe$Iy(4^vY0LlsQU_0t z)=dl==O)~AvzL3ydAf~J593;`3?QN+*m2Ks0sWExlEf)g0zP#MDm`8~>IEba;?=UbQT~S84GD#n@L=c;e%OS z(xRHjdu7$BHcQhoOp-*$4x~62SRaj(yM;~c%~n~ATCY0#?4VTJaZg-NB%H+2Yo)k2 zVeXpJm(yE4varIs)c?(>uT#BEgd%OLg-={B5`za1k;+JrDQ3$;V#PpSfESCdbE>?R zGODFv6uUh7BqDL5V6lh2?NM3awgR6fI@=sp$g>Bj@xK9$zX}V$SNBOEM)L5z)cCJ$ zb!z-auJ}93-wYbYc1Nh|BM(6w4l!HpYoPu@3{>~;B0z)P#&YVw-OpGr5VJpye^-)Y z&KLW_`b&ZJM6a=X-%nEq@`uvjAgq5p_wwAAfY)n-`}!Hq9}290+6(KUCsPOBTZ#V{ z!us~+{+9^r-TA(-UK3c4_rm(%YU;q>{iCt35Y`W$`CgE)PCeTf){g|%CwgH$awT=( zLu%+x2y5=;^-mJkw?E$()?W#%kN3iQ^hSa8uL*1LS0DX#!us-a`-~NQ=Bw*SqWg@* z(2oVy2eyTE=tJtEuNmI&lu&l=bYu~cuoHxvIi6S#u9XJ|IKo@-n+>J_oUa~c96o%X zFylYyWX4Bt1T%gFzue&Ge)|c-_zIvr{p4F(>)QLv9xQ?Kp>3fYdj042f!}Z*b$W1s zus-oyZ|R3oUSK`i1uGyJ9Rx-nHH<$4l%GDf$H-$!eDzT7&Y_&J{m}Y za;IiMxv_}a>DLdjkbv3(O0bMT0!C$UfTOL12Mpw=<#k2=7<+&574x9{O0Acs-f!U7b-j`UMt=*hE5^h4`^A$#e{k=E z-{5!n^lwoNlUz!kn@Pr(5{YCoF`Jl4%p|9i(}`H(ToQlcF+ErY0w2vG9pR!q~rO-+_I`+xM*fmhnFUQ%`SL From d294f50a6e926417bbee7b0600a6f2e1d4f2b897 Mon Sep 17 00:00:00 2001 From: Lachlan Harris Date: Sun, 6 Sep 2026 18:10:44 +1000 Subject: [PATCH 3/4] refc: refactor tests, simplify ci --- .github/dependabot.yaml | 3 +- .github/workflows/ci.yaml | 29 +-- .github/workflows/dependabot-auto-merge.yaml | 30 --- .github/workflows/nightly.yaml | 77 ------ .github/workflows/release.yaml | 5 - .goreleaser.yaml | 18 +- .ko.yaml | 2 +- cmd/check_test.go | 28 --- cmd/targets_test.go | 22 ++ docker/docker-compose.yml | 28 +-- examples/gonetsim-listeners.toml | 8 + internal/capture/session.go | 50 ++-- internal/capture/session_test.go | 102 -------- internal/dnsserver/capture_test.go | 184 ++------------ internal/dnsserver/dns_test.go | 4 +- internal/handler/handler_test.go | 60 +---- internal/handler/lua.go | 3 + internal/handler/luabindings.go | 39 +-- internal/handler/luaconn.go | 17 +- internal/httpserver/capture_test.go | 199 +-------------- internal/httpserver/fakemode.go | 29 +-- internal/httpserver/http_test.go | 244 ++----------------- internal/httpserver/realmode.go | 23 +- internal/httpserver/server.go | 25 ++ internal/listener/listener_test.go | 132 ++-------- internal/listener/service.go | 2 +- internal/listener/tcp.go | 1 - internal/service/manager.go | 4 - internal/service/manager_test.go | 98 -------- internal/state/state.go | 4 +- internal/testutil/testutil.go | 140 +++++++++++ internal/tlsprovider/tls_test.go | 38 +-- 32 files changed, 352 insertions(+), 1296 deletions(-) delete mode 100644 .github/workflows/dependabot-auto-merge.yaml delete mode 100644 .github/workflows/nightly.yaml delete mode 100644 cmd/check_test.go delete mode 100644 internal/service/manager_test.go diff --git a/.github/dependabot.yaml b/.github/dependabot.yaml index d4c9ddf..6de28ec 100644 --- a/.github/dependabot.yaml +++ b/.github/dependabot.yaml @@ -11,7 +11,8 @@ updates: - package-ecosystem: "gomod" directory: "/" # Location of package manifests schedule: - interval: "daily" + interval: "weekly" time: "15:00" + day: "monday" timezone: "Australia/Sydney" labels: ["scope: deps", "priority: medium"] diff --git a/.github/workflows/ci.yaml b/.github/workflows/ci.yaml index 22e52b7..7cd28e8 100644 --- a/.github/workflows/ci.yaml +++ b/.github/workflows/ci.yaml @@ -15,11 +15,8 @@ permissions: jobs: lint: - strategy: - matrix: - os: [ubuntu-latest] name: golangci-lint - runs-on: ${{ matrix.os }} + runs-on: ubuntu-latest steps: - uses: actions/checkout@v7 - name: Install Go @@ -40,52 +37,36 @@ jobs: - name: Checkout code uses: actions/checkout@v7 - name: Install Go - id: setup-go uses: actions/setup-go@v7 with: go-version-file: "go.mod" - - name: Download Go modules - shell: bash - if: ${{ steps.setup-go.outputs.cache-hit != 'true' }} - run: go mod download - name: go build - run: go build + run: go build ./... - name: go test - run: go test -v ./... -race + run: go test ./... -race container: name: container-build + if: github.ref == 'refs/heads/main' runs-on: ubuntu-latest steps: - name: Checkout code uses: actions/checkout@v7 - name: Install Go - id: setup-go uses: actions/setup-go@v7 with: go-version-file: "go.mod" - - name: Download Go modules - shell: bash - if: ${{ steps.setup-go.outputs.cache-hit != 'true' }} - run: go mod download - - name: Install ko run: go install github.com/google/ko@v0.18.1 - name: Set build metadata id: meta shell: bash - env: - PR_NUMBER: ${{ github.event.pull_request.number }} run: | short_sha="${GITHUB_SHA::7}" - if [[ "${GITHUB_EVENT_NAME}" == "pull_request" && -n "${PR_NUMBER}" ]]; then - version="pr-${PR_NUMBER}-${short_sha}" - else - version="${GITHUB_REF_NAME}-${short_sha}" - fi + version="${GITHUB_REF_NAME}-${short_sha}" echo "version=${version}" >> "$GITHUB_OUTPUT" echo "build_date=$(date -u +'%Y-%m-%dT%H:%M:%SZ')" >> "$GITHUB_OUTPUT" diff --git a/.github/workflows/dependabot-auto-merge.yaml b/.github/workflows/dependabot-auto-merge.yaml deleted file mode 100644 index 5c6c0c6..0000000 --- a/.github/workflows/dependabot-auto-merge.yaml +++ /dev/null @@ -1,30 +0,0 @@ -name: Dependabot auto-merge -on: - pull_request: - types: - - opened -permissions: - pull-requests: write - contents: write - repository-projects: write -jobs: - dependabot-automation: - runs-on: ubuntu-latest - if: ${{ github.actor == 'dependabot[bot]' }} - timeout-minutes: 13 - steps: - - name: Dependabot metadata - id: metadata - uses: dependabot/fetch-metadata@v3.1.0 - with: - github-token: ${{ secrets.GITHUB_TOKEN }} - - name: Approve & enable auto-merge for Dependabot PR - if: | - steps.metadata.outputs.update-type == 'version-update:semver-patch' || - steps.metadata.outputs.update-type == 'version-update:semver-minor' - run: | - gh pr merge --auto -s "$PR_URL" - env: - PR_URL: ${{ github.event.pull_request.html_url }} - PR_TITLE: ${{ github.event.pull_request.title }} - GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} diff --git a/.github/workflows/nightly.yaml b/.github/workflows/nightly.yaml deleted file mode 100644 index 35462ad..0000000 --- a/.github/workflows/nightly.yaml +++ /dev/null @@ -1,77 +0,0 @@ -name: Nightly - -on: - schedule: - - cron: '0 2 * * *' # Runs at 02:00 UTC every night - workflow_dispatch: - -permissions: - contents: read - packages: write - - -jobs: - check-and-build: - runs-on: ubuntu-latest - steps: - - name: Check for new commits in the last 24h - id: check - uses: adriangl/check-new-commits-action@v2 - with: - token: ${{ secrets.GITHUB_TOKEN }} - seconds: 86400 # 24 hours - - - name: Checkout code - if: steps.check.outputs.has-new-commits == 'true' - uses: actions/checkout@v7 - - - name: Install Go - if: steps.check.outputs.has-new-commits == 'true' - id: setup-go - uses: actions/setup-go@v7 - with: - go-version-file: "go.mod" - - - name: Download Go modules - if: ${{ steps.check.outputs.has-new-commits == 'true' && steps.setup-go.outputs.cache-hit != 'true' }} - shell: bash - run: go mod download - - - name: Install ko - if: steps.check.outputs.has-new-commits == 'true' - run: go install github.com/google/ko@v0.18.1 - - - name: Set build metadata - if: steps.check.outputs.has-new-commits == 'true' - id: meta - shell: bash - run: | - date="$(date -u +'%Y%m%d')" - echo "date=${date}" >> "$GITHUB_OUTPUT" - echo "version=nightly-${date}" >> "$GITHUB_OUTPUT" - echo "build_date=$(date -u +'%Y-%m-%dT%H:%M:%SZ')" >> "$GITHUB_OUTPUT" - - - name: Log in to GHCR - if: steps.check.outputs.has-new-commits == 'true' - uses: docker/login-action@v4 - with: - registry: ghcr.io - username: ${{ github.actor }} - password: ${{ secrets.GITHUB_TOKEN }} - - - name: Build and push image (ko) - if: steps.check.outputs.has-new-commits == 'true' - env: - KO_DOCKER_REPO: ghcr.io/${{ github.repository_owner }}/gonetsim - VERSION: ${{ steps.meta.outputs.version }} - REVISION: ${{ github.sha }} - BUILD_DATE: ${{ steps.meta.outputs.build_date }} - shell: bash - run: | - ko build \ - --bare \ - --platform=linux/amd64,linux/arm64 \ - --sbom=none \ - --image-user=0:0 \ - --tags=nightly,nightly-${{ steps.meta.outputs.date }} \ - . diff --git a/.github/workflows/release.yaml b/.github/workflows/release.yaml index cd3f550..c861684 100644 --- a/.github/workflows/release.yaml +++ b/.github/workflows/release.yaml @@ -22,14 +22,9 @@ jobs: with: fetch-depth: 0 - name: Install Go - id: setup-go uses: actions/setup-go@v7 with: go-version-file: "go.mod" - - name: Download Go modules - shell: bash - if: ${{ steps.setup-go.outputs.cache-hit != 'true' }} - run: go mod download - name: Install ko run: go install github.com/google/ko@v0.18.1 - name: Log in to GHCR diff --git a/.goreleaser.yaml b/.goreleaser.yaml index c8f2a6b..f3f7d22 100644 --- a/.goreleaser.yaml +++ b/.goreleaser.yaml @@ -98,20 +98,4 @@ release: - [Documentation](https://gonetsim.lachlanharris.au/) - [Report Issues](https://github.com/{{ .Env.GITHUB_REPOSITORY }}/issues) - - [Discussions](https://github.com/{{ .Env.GITHUB_REPOSITORY }}/discussions) - -kos: - - id: gonetsim - build: gonetsim - repositories: - - ghcr.io/{{ .Env.GITHUB_REPOSITORY_OWNER }}/gonetsim - tags: - - "{{ .Version }}" - - "{{ .Major }}.{{ .Minor }}" - - latest - bare: true - platforms: - - linux/amd64 - - linux/arm64 - base_image: gcr.io/distroless/base-debian12 - user: "0:0" + - [Discussions](https://github.com/{{ .Env.GITHUB_REPOSITORY }}/discussions) \ No newline at end of file diff --git a/.ko.yaml b/.ko.yaml index 6de211e..e8f03d2 100644 --- a/.ko.yaml +++ b/.ko.yaml @@ -17,4 +17,4 @@ builds: - -w - -X github.com/lachlanharrisdev/gonetsim/cmd.Version={{.Env.VERSION}} - -X github.com/lachlanharrisdev/gonetsim/cmd.Revision={{.Env.REVISION}} - - -X github.com/lachlanharrisdev/gonetsim/cmd.BuildDate={{.Env.BUILD_DATE}} + - -X github.com/lachlanharrisdev/gonetsim/cmd.Date={{.Env.BUILD_DATE}} diff --git a/cmd/check_test.go b/cmd/check_test.go deleted file mode 100644 index 0142b25..0000000 --- a/cmd/check_test.go +++ /dev/null @@ -1,28 +0,0 @@ -package cmd - -import ( - "os" - "path/filepath" - "testing" - - "github.com/lachlanharrisdev/gonetsim/internal/capture" -) - -func TestCheckRunDir(t *testing.T) { - dir := t.TempDir() - t.Setenv("XDG_DATA_HOME", dir) - - if err := checkRunDir(); err != nil { - t.Fatalf("checkRunDir: %v", err) - } - runs, err := capture.DefaultRunsDir() - if err != nil { - t.Fatalf("DefaultRunsDir: %v", err) - } - if st, err := os.Stat(runs); err != nil || !st.IsDir() { - t.Fatalf("expected runs dir to exist: %v", err) - } - if filepath.Dir(runs) != filepath.Join(dir, "gonetsim") { - t.Fatalf("runs dir = %q, want it under %q", runs, dir) - } -} diff --git a/cmd/targets_test.go b/cmd/targets_test.go index 4183ce8..69e2793 100644 --- a/cmd/targets_test.go +++ b/cmd/targets_test.go @@ -2,9 +2,12 @@ package cmd import ( "log/slog" + "os" + "path/filepath" "strings" "testing" + "github.com/lachlanharrisdev/gonetsim/internal/capture" appconfig "github.com/lachlanharrisdev/gonetsim/internal/config" "github.com/lachlanharrisdev/gonetsim/internal/state" "github.com/lachlanharrisdev/gonetsim/internal/testutil" @@ -243,3 +246,22 @@ func testResolve(t *testing.T, cfg *appconfig.Config, args []string, opts runOpt } return resolved } + +func TestCheckRunDir(t *testing.T) { + dir := t.TempDir() + t.Setenv("XDG_DATA_HOME", dir) + + if err := checkRunDir(); err != nil { + t.Fatalf("checkRunDir: %v", err) + } + runs, err := capture.DefaultRunsDir() + if err != nil { + t.Fatalf("DefaultRunsDir: %v", err) + } + if st, err := os.Stat(runs); err != nil || !st.IsDir() { + t.Fatalf("expected runs dir to exist: %v", err) + } + if filepath.Dir(runs) != filepath.Join(dir, "gonetsim") { + t.Fatalf("runs dir = %q, want it under %q", runs, dir) + } +} diff --git a/docker/docker-compose.yml b/docker/docker-compose.yml index 8cdb47b..566e000 100644 --- a/docker/docker-compose.yml +++ b/docker/docker-compose.yml @@ -1,38 +1,14 @@ services: gonetsim: - # for local, build the image locally with ko, then run it with compose: - # go install github.com/google/ko@v0.18.1 - # cd .. - # KO_DOCKER_REPO=gonetsim ko build --local --bare --tags=dev . - # then set the image to: - # image: gonetsim:dev + # build locally with `KO_DOCKER_REPO=gonetsim ko build --local --bare --tags=dev .` + # then set image: gonetsim:dev image: ghcr.io/lachlanharrisdev/gonetsim:latest - - # publish the default non-privileged ports - # to publish standard ports on the host instead, swap to: - # - "53:53/udp" - # - "53:53/tcp" - # - "80:80" - # - "443:443" - # note that updating the configuration file won't change the affected host ports, unlike the standalone binary ports: - "5353:53/udp" - "5353:53/tcp" - "8080:80" - "8443:443" - - # by default GoNetSim will generate a config file on first start. - # if you want to pin the config to /etc/gonetsim/gonetsim.toml, mount it there: - # volumes: - # - ../gonetsim.toml:/etc/gonetsim/gonetsim.toml:ro - - # custom listener Lua scripts are resolved relative to the config file's - # directory, so mount them alongside it: # volumes: # - ../gonetsim.toml:/etc/gonetsim/gonetsim.toml:ro # - ../handlers:/etc/gonetsim/handlers:ro - - # the run capture is written inside the container's data directory; - # mount a volume there to keep it, or pass --output to choose a path: - # volumes: # - ./captures:/root/.local/share/gonetsim/runs diff --git a/examples/gonetsim-listeners.toml b/examples/gonetsim-listeners.toml index b6bcfa0..10d5f10 100644 --- a/examples/gonetsim-listeners.toml +++ b/examples/gonetsim-listeners.toml @@ -5,6 +5,8 @@ # - ftp : fake FTP server (lua:handlers/ftp.lua) # - echo : TCP echo service (builtin:echo) # - sink : UDP discard sink (builtin:sink) +# A fifth example (smtp, lua:handlers/smtp.lua) is commented out below +# uncomment to try the AUTH/state demo on :2525 # # Run from the repository root with: # gonetsim --config examples/gonetsim-listeners.toml @@ -59,3 +61,9 @@ name = "sink" type = "udp" listen = ":9999" handler = "builtin:sink" + +# [[listeners]] +# name = "smtp" +# type = "tcp" +# listen = ":2525" +# handler = "lua:handlers/smtp.lua" diff --git a/internal/capture/session.go b/internal/capture/session.go index 6de4452..314317b 100644 --- a/internal/capture/session.go +++ b/internal/capture/session.go @@ -22,8 +22,11 @@ type Session struct { local netip.AddrPort remote netip.AddrPort + // the pcapng holds no real handshake, so thefirst Write emits SYN, SYN-ACK, + // then data with synthetic seq/ack numbers tracked in clientSeq/serverSeq + // Close emits FIN/FIN-ACK synSent bool - pending string + pending string // comment attached to the next emitted frame clientSeq uint32 serverSeq uint32 @@ -201,26 +204,28 @@ func (s *Session) build(data []byte, src, dst netip.AddrPort, kind transportKind switch kind { case isTCP: tcp := &layers.TCP{SrcPort: layers.TCPPort(src.Port()), DstPort: layers.TCPPort(dst.Port())} - network = tcpIPLayer(src, dst, layers.IPProtocolTCP, tcp) + network = ipLayer(src, dst, layers.IPProtocolTCP, tcp) transport = tcp default: udp := &layers.UDP{SrcPort: layers.UDPPort(src.Port()), DstPort: layers.UDPPort(dst.Port())} - network = udpIPLayer(src, dst, udp) + network = ipLayer(src, dst, layers.IPProtocolUDP, udp) transport = udp } return s.serialize(data, src, network, transport) } func (s *Session) buildTCPData(data []byte, fromClient bool, seq, ack uint32) []byte { - src, dst := s.endpoints(fromClient) - tcp := newTCPLayer(src, dst, seq, ack, false, true, false, len(data) > 0) - return s.buildWithTCP(data, src, dst, tcp) + return s.buildTCP(data, fromClient, false, true, false, len(data) > 0, seq, ack) } func (s *Session) buildTCPControl(fromClient, syn, ackFlag, fin bool, seq, ackNum uint32) []byte { + return s.buildTCP(nil, fromClient, syn, ackFlag, fin, false, seq, ackNum) +} + +func (s *Session) buildTCP(data []byte, fromClient, syn, ackFlag, fin, psh bool, seq, ack uint32) []byte { src, dst := s.endpoints(fromClient) - tcp := newTCPLayer(src, dst, seq, ackNum, syn, ackFlag, fin, false) - return s.buildWithTCP(nil, src, dst, tcp) + tcp := newTCPLayer(src, dst, seq, ack, syn, ackFlag, fin, psh) + return s.serialize(data, src, ipLayer(src, dst, layers.IPProtocolTCP, tcp), tcp) } func newTCPLayer(src, dst netip.AddrPort, seq, ack uint32, syn, ackFlag, fin, psh bool) *layers.TCP { @@ -237,29 +242,22 @@ func newTCPLayer(src, dst netip.AddrPort, seq, ack uint32, syn, ackFlag, fin, ps } } -func (s *Session) buildWithTCP(data []byte, src, dst netip.AddrPort, tcp *layers.TCP) []byte { - return s.serialize(data, src, tcpIPLayer(src, dst, layers.IPProtocolTCP, tcp), tcp) -} - -func tcpIPLayer(src, dst netip.AddrPort, proto layers.IPProtocol, tcp *layers.TCP) gopacket.SerializableLayer { +func ipLayer(src, dst netip.AddrPort, proto layers.IPProtocol, transport gopacket.SerializableLayer) gopacket.SerializableLayer { + setChecksum := func(ip gopacket.NetworkLayer) { + switch t := transport.(type) { + case *layers.TCP: + _ = t.SetNetworkLayerForChecksum(ip) + case *layers.UDP: + _ = t.SetNetworkLayerForChecksum(ip) + } + } if src.Addr().Is4() { ip := &layers.IPv4{Version: 4, TTL: 64, Protocol: proto, SrcIP: src.Addr().AsSlice(), DstIP: dst.Addr().AsSlice()} - _ = tcp.SetNetworkLayerForChecksum(ip) + setChecksum(ip) return ip } ip := &layers.IPv6{Version: 6, HopLimit: 64, NextHeader: proto, SrcIP: src.Addr().AsSlice(), DstIP: dst.Addr().AsSlice()} - _ = tcp.SetNetworkLayerForChecksum(ip) - return ip -} - -func udpIPLayer(src, dst netip.AddrPort, udp *layers.UDP) gopacket.SerializableLayer { - if src.Addr().Is4() { - ip := &layers.IPv4{Version: 4, TTL: 64, Protocol: layers.IPProtocolUDP, SrcIP: src.Addr().AsSlice(), DstIP: dst.Addr().AsSlice()} - _ = udp.SetNetworkLayerForChecksum(ip) - return ip - } - ip := &layers.IPv6{Version: 6, HopLimit: 64, NextHeader: layers.IPProtocolUDP, SrcIP: src.Addr().AsSlice(), DstIP: dst.Addr().AsSlice()} - _ = udp.SetNetworkLayerForChecksum(ip) + setChecksum(ip) return ip } diff --git a/internal/capture/session_test.go b/internal/capture/session_test.go index 2a80799..e84e6bd 100644 --- a/internal/capture/session_test.go +++ b/internal/capture/session_test.go @@ -10,7 +10,6 @@ package capture import ( "bytes" - "encoding/binary" "io" "net/netip" "os" @@ -229,20 +228,6 @@ func TestSessionComment(t *testing.T) { } } -func TestSessionEmpty(t *testing.T) { - local := netip.MustParseAddrPort("127.0.0.1:8080") - remote := netip.MustParseAddrPort("10.0.0.5:40000") - run, path := testRun(t) - - ses := testSession(t, run, "tcp", local, remote) - if err := ses.Close(); err != nil { - t.Fatalf("Close: %v", err) - } - if pkts := readFrames(t, path); len(pkts) != 0 { - t.Fatalf("expected empty capture, got %d packets", len(pkts)) - } -} - func TestRun(t *testing.T) { t.Run("one file holds many flows", func(t *testing.T) { run, path := testRun(t) @@ -310,53 +295,6 @@ func TestRun(t *testing.T) { t.Fatalf("expected parent dir to be created: %v", err) } }) - - t.Run("run id format", func(t *testing.T) { - id := NewRunID() - if len(id) != 20 || id[8] != '-' || id[15] != '-' { - t.Fatalf("unexpected run id %q", id) - } - }) - - t.Run("empty run inspects cleanly", func(t *testing.T) { - run, path := testRun(t) - if err := run.Close(); err != nil { - t.Fatalf("Close: %v", err) - } - info, err := Inspect(path) - if err != nil { - t.Fatalf("Inspect: %v", err) - } - if info.Packets != 0 { - t.Fatalf("Packets = %d, want 0", info.Packets) - } - }) - - t.Run("nil run is a no-op", func(t *testing.T) { - var run *Run - if run.Path() != "" { - t.Fatalf("expected empty path from nil run") - } - if packets, first, last := run.Stats(); packets != 0 || !first.IsZero() || !last.IsZero() { - t.Fatalf("expected zero stats from nil run") - } - if err := run.Close(); err != nil { - t.Fatalf("Close on nil run: %v", err) - } - if iface, err := run.NewInterface("x"); err != nil || iface != 0 { - t.Fatalf("expected zero interface from nil run, got %d, %v", iface, err) - } - ses, err := run.NewSession("tcp", - netip.MustParseAddrPort("5.6.7.8:9"), netip.MustParseAddrPort("1.2.3.4:5"), 0) - if err != nil || ses != nil { - t.Fatalf("expected nil session from nil run, got %v, %v", ses, err) - } - ses.Comment("ignored") - _ = ses.Write([]byte("data"), true) - if err := ses.Close(); err != nil { - t.Fatalf("Close on nil session: %v", err) - } - }) } func TestInspect(t *testing.T) { @@ -409,43 +347,3 @@ func TestInspect(t *testing.T) { t.Fatalf("expected legacy pcap error, got %v", err) } } - -func TestBlockLengthsMatch(t *testing.T) { - local := netip.MustParseAddrPort("127.0.0.1:8080") - remote := netip.MustParseAddrPort("10.0.0.5:40000") - run, path := testRun(t) - - ses := testSession(t, run, "tcp", local, remote) - ses.Comment("greeting") - if err := ses.Write([]byte("hello"), true); err != nil { - t.Fatalf("Write: %v", err) - } - if err := ses.Close(); err != nil { - t.Fatalf("Close: %v", err) - } - - raw, err := os.ReadFile(path) - if err != nil { - t.Fatalf("ReadFile: %v", err) - } - off := 0 - blocks := 0 - for off < len(raw) { - if len(raw)-off < 12 { - t.Fatalf("block %d at offset %d: truncated header", blocks, off) - } - blen := int(binary.LittleEndian.Uint32(raw[off+4 : off+8])) - if blen < 12 || off+blen > len(raw) { - t.Fatalf("block %d at offset %d: bad length %d (file %d)", blocks, off, blen, len(raw)) - } - trail := binary.LittleEndian.Uint32(raw[off+blen-4 : off+blen]) - if int(trail) != blen { - t.Fatalf("block %d at offset %d: lengths %d and %d don't match", blocks, off, blen, trail) - } - off += blen - blocks++ - } - if blocks == 0 { - t.Fatalf("no blocks found") - } -} diff --git a/internal/dnsserver/capture_test.go b/internal/dnsserver/capture_test.go index c6c0d1b..914f19d 100644 --- a/internal/dnsserver/capture_test.go +++ b/internal/dnsserver/capture_test.go @@ -10,18 +10,11 @@ package dnsserver import ( "context" - "fmt" - "net" "net/netip" - "os" - "path/filepath" "strings" "testing" "time" - "github.com/google/gopacket" - "github.com/google/gopacket/layers" - "github.com/google/gopacket/pcapgo" "github.com/miekg/dns" "github.com/lachlanharrisdev/gonetsim/internal/capture" @@ -29,80 +22,36 @@ import ( "github.com/lachlanharrisdev/gonetsim/internal/testutil" ) -func testRun(t *testing.T) (*capture.Run, string) { - t.Helper() - path := filepath.Join(t.TempDir(), "run.pcapng") - run, err := capture.NewRun(path) - if err != nil { - t.Fatalf("NewRun: %v", err) - } - t.Cleanup(func() { _ = run.Close() }) - return run, path -} - -func TestService_CapturesUDP(t *testing.T) { - conf := baseCaptureConfig(t, "udp") - conf.Capture = true - run, path := testRun(t) - - svc, errCh := startDNSService(t, conf, run) - - query := newAQuery() - client := &dns.Client{Net: "udp", Timeout: 1 * time.Second} - _, _, err := retryExchange(t, client, conf.Addr, query) - if err != nil { - t.Fatalf("exchange: %v", err) - } - - waitDNSCapture(t, path, "example") - waitTransportPayloads(t, path, func(s string) bool { - return strings.Count(s, "example") >= 2 // query + response - }) - - svc.Stop(context.Background()) //nolint:errcheck,gosec - discardStartErr(t, errCh) -} - -func TestService_CapturesTCP(t *testing.T) { - conf := baseCaptureConfig(t, "tcp") - conf.Capture = true - run, path := testRun(t) - - svc, errCh := startDNSService(t, conf, run) - - query := newAQuery() - client := &dns.Client{Net: "tcp", Timeout: 1 * time.Second} - _, _, err := retryExchange(t, client, conf.Addr, query) - if err != nil { - t.Fatalf("exchange: %v", err) +func TestService_Captures(t *testing.T) { + for _, network := range []string{"udp", "tcp"} { + t.Run(network, func(t *testing.T) { + conf := baseCaptureConfig(t, network) + conf.Capture = true + run, path := testutil.NewPcapRun(t) + + svc, errCh := startDNSService(t, conf, run) + + query := newAQuery() + client := &dns.Client{Net: network, Timeout: 1 * time.Second} + if _, _, err := testutil.RetryDNSExchange(t, client, conf.Addr, query); err != nil { + t.Fatalf("exchange: %v", err) + } + + testutil.WaitForPayloadContains(t, path, "example", 3*time.Second) + testutil.WaitForPayload(t, path, 3*time.Second, func(s string) bool { + return strings.Count(s, "example") >= 2 // query + response + }) + + _ = svc.Stop(context.Background()) + testutil.DiscardServiceStartErr(t, errCh) + }) } - - waitDNSCapture(t, path, "example") - waitTransportPayloads(t, path, func(s string) bool { - return strings.Count(s, "example") >= 2 // query + response - }) - - svc.Stop(context.Background()) //nolint:errcheck,gosec - discardStartErr(t, errCh) } func baseCaptureConfig(t *testing.T, network string) Config { t.Helper() - listener, err := net.Listen("tcp", "127.0.0.1:0") - if err != nil { - t.Fatalf("Listen: %v", err) - } - port := listener.Addr().(*net.TCPAddr).Port - _ = listener.Close() - - pc, err := net.ListenPacket("udp", "127.0.0.1:0") - if err != nil { - t.Fatalf("ListenPacket: %v", err) - } - _ = pc.Close() - return Config{ - Addr: fmt.Sprintf("127.0.0.1:%d", port), + Addr: testutil.FreePort(t, "tcp"), Net: network, SinkholeIPv4: netip.MustParseAddr("203.0.113.10"), SinkholeIPv6: netip.MustParseAddr("2001:db8::10"), @@ -128,88 +77,3 @@ func startDNSService(t *testing.T, conf Config, run *capture.Run) (service.Servi go func() { errCh <- svc.Start(context.Background()) }() return svc, errCh } - -func retryExchange(t *testing.T, client *dns.Client, addr string, m *dns.Msg) (*dns.Msg, time.Duration, error) { - t.Helper() - deadline := time.Now().Add(3 * time.Second) - var lastErr error - var lastRTT time.Duration - for time.Now().Before(deadline) { - resp, rtt, err := client.Exchange(m, addr) - if err == nil && resp != nil { - return resp, rtt, nil - } - lastErr, lastRTT = err, rtt - time.Sleep(20 * time.Millisecond) - } - return nil, lastRTT, lastErr -} - -func discardStartErr(t *testing.T, errCh <-chan error) { - t.Helper() - select { - case err := <-errCh: - if err != nil { - t.Fatalf("service.Start returned error: %v", err) - } - case <-time.After(3 * time.Second): - t.Fatalf("service.Start never returned") - } -} - -// waitDNSCapture waits for the run capture's transport payloads to contain -// want. -func waitDNSCapture(t *testing.T, path, want string) { - t.Helper() - deadline := time.Now().Add(3 * time.Second) - for time.Now().Before(deadline) { - if payloads, err := transportPayloads(path); err == nil && strings.Contains(payloads, want) { - return - } - time.Sleep(20 * time.Millisecond) - } - t.Fatalf("run capture %s never contained %q", path, want) -} - -// waitTransportPayloads polls until the transport payloads of the capture -// satisfy cond (tolerating the async flush on connection teardown). -func waitTransportPayloads(t *testing.T, path string, cond func(string) bool) { - t.Helper() - deadline := time.Now().Add(3 * time.Second) - for time.Now().Before(deadline) { - joined, err := transportPayloads(path) - if err == nil && cond(joined) { - return - } - time.Sleep(20 * time.Millisecond) - } - joined, _ := transportPayloads(path) - t.Fatalf("capture %s never satisfied predicate (payloads %q)", path, joined) -} - -// transportPayloads concatenates UDP/TCP payloads from a pcapng file. -func transportPayloads(path string) (string, error) { - f, err := os.Open(path) - if err != nil { - return "", err - } - defer f.Close() - r, err := pcapgo.NewNgReader(f, pcapgo.DefaultNgReaderOptions) - if err != nil { - return "", err - } - var sb strings.Builder - for { - data, _, err := r.ReadPacketData() - if err != nil { - break - } - pkt := gopacket.NewPacket(data, layers.LinkTypeEthernet, gopacket.Default) - if u, ok := pkt.Layer(layers.LayerTypeUDP).(*layers.UDP); ok { - sb.Write(u.Payload) - } else if tcp, ok := pkt.Layer(layers.LayerTypeTCP).(*layers.TCP); ok { - sb.Write(tcp.Payload) - } - } - return sb.String(), nil -} diff --git a/internal/dnsserver/dns_test.go b/internal/dnsserver/dns_test.go index 67736b5..2621b4e 100644 --- a/internal/dnsserver/dns_test.go +++ b/internal/dnsserver/dns_test.go @@ -165,10 +165,10 @@ func TestRecordTypes(t *testing.T) { {"SOA", "example.com.", dns.TypeSOA, checkSOA}, {"CAA", "example.com.", dns.TypeCAA, checkCAA}, } + client, addr, conf, teardown := queryTestsHelper(t) + defer teardown() for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { - client, addr, conf, teardown := queryTestsHelper(t) - defer teardown() resp := exchange(t, client, addr, tc.qname, tc.qtype) if len(resp.Answer) != 1 { t.Fatalf("expected 1 answer, got %d", len(resp.Answer)) diff --git a/internal/handler/handler_test.go b/internal/handler/handler_test.go index 4920870..6b59d37 100644 --- a/internal/handler/handler_test.go +++ b/internal/handler/handler_test.go @@ -12,17 +12,14 @@ import ( "io" "log/slog" "net" - "net/netip" - "path/filepath" "strings" "testing" - "github.com/lachlanharrisdev/gonetsim/internal/capture" "github.com/lachlanharrisdev/gonetsim/internal/state" "github.com/lachlanharrisdev/gonetsim/internal/testutil" ) -func discardLogger() *slog.Logger { +func testLogger() *slog.Logger { return testutil.Logger() } @@ -56,7 +53,7 @@ func roundtrip(t *testing.T, client net.Conn, payload, reply string) { func TestBuiltins(t *testing.T) { t.Run("tcp echo", func(t *testing.T) { - client, done := servePipe(t, EchoHandler{}, Env{Logger: discardLogger()}) + client, done := servePipe(t, EchoHandler{}, Env{Logger: testLogger()}) roundtrip(t, client, "abc", "abc") _ = client.Close() if err := <-done; err != nil { @@ -66,7 +63,7 @@ func TestBuiltins(t *testing.T) { t.Run("udp echo", func(t *testing.T) { addr, _ := net.ResolveUDPAddr("udp", "203.0.113.10:53") - reply, err := EchoHandler{}.HandleUDP(t.Context(), []byte("query"), addr, Env{Logger: discardLogger()}) + reply, err := EchoHandler{}.HandleUDP(t.Context(), []byte("query"), addr, Env{Logger: testLogger()}) if err != nil || string(reply) != "query" { t.Fatalf("udp echo: %v %q", err, reply) } @@ -87,7 +84,7 @@ func TestLuaHandler(t *testing.T) { if err != nil { t.Fatalf("NewLua: %v", err) } - client, done := servePipe(t, h, Env{Logger: discardLogger()}) + client, done := servePipe(t, h, Env{Logger: testLogger()}) roundtrip(t, client, "hello\nworld\n", "echo: hello\necho: world\n") _ = client.Close() if err := <-done; err != nil { @@ -101,56 +98,15 @@ func TestLuaHandler(t *testing.T) { if err != nil { t.Fatalf("NewLua: %v", err) } - reply, err := h.HandleUDP(t.Context(), []byte("ping"), remote, Env{Logger: discardLogger()}) + reply, err := h.HandleUDP(t.Context(), []byte("ping"), remote, Env{Logger: testLogger()}) if err != nil || string(reply) != "pong" { t.Fatalf("ping: %v %q", err, reply) } - reply, err = h.HandleUDP(t.Context(), []byte("other"), remote, Env{Logger: discardLogger()}) + reply, err = h.HandleUDP(t.Context(), []byte("other"), remote, Env{Logger: testLogger()}) if err != nil || reply != nil { t.Fatalf("silent: %v %q", err, reply) } }) - - t.Run("capture comment", func(t *testing.T) { - h, err := NewLua("testdata/comment.lua", nil) - if err != nil { - t.Fatalf("NewLua: %v", err) - } - c1, c2 := net.Pipe() - defer c2.Close() - done := make(chan error, 1) - go func() { - done <- h.HandleTCP(t.Context(), c2, Env{Logger: discardLogger()}) - }() - - run, err := capture.NewRun(filepath.Join(t.TempDir(), "run.pcapng")) - if err != nil { - t.Fatalf("NewRun: %v", err) - } - defer func() { _ = run.Close() }() - iface, err := run.NewInterface("test") - if err != nil { - t.Fatalf("NewInterface: %v", err) - } - ses, err := run.NewSession("tcp", - netip.MustParseAddrPort("127.0.0.1:9"), netip.MustParseAddrPort("203.0.113.10:1"), iface) - if err != nil { - t.Fatalf("NewSession: %v", err) - } - ses.Comment("client said hello") - _ = ses.Write([]byte("hello"), true) - if err := ses.Close(); err != nil { - t.Fatalf("Close: %v", err) - } - - _, _ = c1.Write([]byte("hello")) - if err := c1.Close(); err != nil { - t.Fatalf("Close: %v", err) - } - if err := <-done; err != nil { - t.Fatalf("HandleTCP: %v", err) - } - }) } func TestLuaState(t *testing.T) { @@ -158,7 +114,7 @@ func TestLuaState(t *testing.T) { if err != nil { t.Fatalf("NewLua: %v", err) } - env := Env{Logger: discardLogger(), Global: state.NewStore(state.NewBudget(state.DefaultTotalLimit))} + env := Env{Logger: testLogger(), Global: state.NewStore(state.NewBudget(state.DefaultTotalLimit))} for i, want := range []string{"1|conn|yes", "2|conn|yes"} { client, done := servePipe(t, h, env) buf := make([]byte, len(want)) @@ -180,7 +136,7 @@ func TestSandboxGlobals(t *testing.T) { if err != nil { t.Fatalf("NewLua: %v", err) } - client, done := servePipe(t, h, Env{Logger: discardLogger()}) + client, done := servePipe(t, h, Env{Logger: testLogger()}) buf := make([]byte, 1024) n, err := client.Read(buf) if err != nil { diff --git a/internal/handler/lua.go b/internal/handler/lua.go index 040bbe2..ed3601b 100644 --- a/internal/handler/lua.go +++ b/internal/handler/lua.go @@ -107,6 +107,9 @@ func (h *LuaHandler) run(L *lua.LState, entry string, nret int, args ...lua.LVal if err := L.CallByParam(lua.P{Fn: fn, NRet: nret, Protect: true}, args...); err != nil { return nil, err } + if nret == 0 { + return nil, nil + } vals := make([]lua.LValue, nret) for i := range vals { vals[i] = L.Get(-nret + i) diff --git a/internal/handler/luabindings.go b/internal/handler/luabindings.go index 7020b11..61e18f3 100644 --- a/internal/handler/luabindings.go +++ b/internal/handler/luabindings.go @@ -180,30 +180,15 @@ func registerConn(L *lua.LState, lc *luaConn, ctx context.Context, env Env, conn })) L.SetField(conn, "remote_ip", L.NewFunction(func(L *lua.LState) int { - if tcp, ok := lc.RemoteAddr().(*net.TCPAddr); ok { - L.Push(lua.LString(tcp.IP.String())) - } else { - L.Push(lua.LNil) - } - return 1 + return pushTCPIP(L, lc.RemoteAddr()) })) L.SetField(conn, "remote_port", L.NewFunction(func(L *lua.LState) int { - if tcp, ok := lc.RemoteAddr().(*net.TCPAddr); ok { - L.Push(lua.LNumber(tcp.Port)) - } else { - L.Push(lua.LNil) - } - return 1 + return pushTCPPort(L, lc.RemoteAddr()) })) L.SetField(conn, "local_port", L.NewFunction(func(L *lua.LState) int { - if tcp, ok := lc.LocalAddr().(*net.TCPAddr); ok { - L.Push(lua.LNumber(tcp.Port)) - } else { - L.Push(lua.LNil) - } - return 1 + return pushTCPPort(L, lc.LocalAddr()) })) L.SetField(conn, "sni", L.NewFunction(func(L *lua.LState) int { @@ -241,3 +226,21 @@ func pushResult(L *lua.LState, v lua.LValue, err error) int { L.Push(v) return 1 } + +func pushTCPIP(L *lua.LState, addr net.Addr) int { + if tcp, ok := addr.(*net.TCPAddr); ok { + L.Push(lua.LString(tcp.IP.String())) + } else { + L.Push(lua.LNil) + } + return 1 +} + +func pushTCPPort(L *lua.LState, addr net.Addr) int { + if tcp, ok := addr.(*net.TCPAddr); ok { + L.Push(lua.LNumber(tcp.Port)) + } else { + L.Push(lua.LNil) + } + return 1 +} diff --git a/internal/handler/luaconn.go b/internal/handler/luaconn.go index 1fdeb57..9d98428 100644 --- a/internal/handler/luaconn.go +++ b/internal/handler/luaconn.go @@ -89,10 +89,7 @@ func (lc *luaConn) readLine() (lua.LValue, error) { continue } if errors.Is(err, io.EOF) { - if sb.Len() > 0 { - return lua.LString(sb.String()), nil - } - return lua.LNil, nil + return eofValue(sb.String()), nil } return nil, err } @@ -112,12 +109,16 @@ func (lc *luaConn) readUntil(delim []byte) (lua.LValue, error) { buf = append(buf, tmp[:n]...) if err != nil { if errors.Is(err, io.EOF) { - if len(buf) > 0 { - return lua.LString(buf), nil - } - return lua.LNil, nil + return eofValue(string(buf)), nil } return nil, err } } } + +func eofValue(s string) lua.LValue { + if s != "" { + return lua.LString(s) + } + return lua.LNil +} diff --git a/internal/httpserver/capture_test.go b/internal/httpserver/capture_test.go index d036a9f..c445ae4 100644 --- a/internal/httpserver/capture_test.go +++ b/internal/httpserver/capture_test.go @@ -10,49 +10,28 @@ package httpserver import ( "context" - "crypto/tls" "io" - "net" "net/http" - "os" - "path/filepath" - "strings" "testing" "time" - "github.com/google/gopacket" - "github.com/google/gopacket/layers" - "github.com/google/gopacket/pcapgo" - "github.com/lachlanharrisdev/gonetsim/internal/capture" "github.com/lachlanharrisdev/gonetsim/internal/service" "github.com/lachlanharrisdev/gonetsim/internal/testutil" - "github.com/lachlanharrisdev/gonetsim/internal/tlsprovider" ) -func testRun(t *testing.T) (*capture.Run, string) { - t.Helper() - path := filepath.Join(t.TempDir(), "run.pcapng") - run, err := capture.NewRun(path) - if err != nil { - t.Fatalf("NewRun: %v", err) - } - t.Cleanup(func() { _ = run.Close() }) - return run, path -} - func TestService_CapturesHTTP(t *testing.T) { conf := Config{ - Addr: freeTCPAddr(t), + Addr: testutil.FreeTCPAddr(t), StatusCode: http.StatusOK, Mode: "fake", Capture: true, } - run, path := testRun(t) + run, path := testutil.NewPcapRun(t) svc, errCh := startHTTPService(t, conf, run) get := func(url string) *http.Response { - _, resp := retryingGet(t, http.DefaultClient, url) + _, resp := testutil.RetryGet(t, http.DefaultClient, url) return resp } get("http://" + conf.Addr + "/warmup") @@ -60,84 +39,11 @@ func TestService_CapturesHTTP(t *testing.T) { _, _ = io.Copy(io.Discard, resp.Body) _ = resp.Body.Close() - waitHTTPCapture(t, path, "GET /hello") - waitTCPPayloadsContain(t, path, "HTTP/1.1 200") - - svc.Stop(context.Background()) //nolint:errcheck,gosec - discardStartErr(t, errCh) -} + testutil.WaitForPayloadContains(t, path, "GET /hello", 3*time.Second) + testutil.WaitForPayloadContains(t, path, "HTTP/1.1 200", 3*time.Second) -func TestService_CapturesHTTPS(t *testing.T) { - dir := t.TempDir() - certPEM, keyPEM, _, err := tlsprovider.GenerateSelfSignedWithCA(tlsprovider.SelfSignedOptions{DNSNames: []string{"localhost"}}) - if err != nil { - t.Fatalf("GenerateSelfSignedWithCA: %v", err) - } - certPath := filepath.Join(dir, "cert.pem") - keyPath := filepath.Join(dir, "key.pem") - if err := os.WriteFile(certPath, certPEM, 0o644); err != nil { - t.Fatalf("WriteFile: %v", err) - } - if err := os.WriteFile(keyPath, keyPEM, 0o600); err != nil { - t.Fatalf("WriteFile: %v", err) - } - - conf := Config{ - Addr: freeTCPAddr(t), - StatusCode: http.StatusOK, - Mode: "fake", - TLS: &tlsprovider.Config{CertFile: certPath, KeyFile: keyPath}, - Capture: true, - } - run, path := testRun(t) - svc, errCh := startHTTPService(t, conf, run) - - client := &http.Client{ - Timeout: 3 * time.Second, - Transport: &http.Transport{ - TLSClientConfig: &tls.Config{InsecureSkipVerify: true}, - }, - } - _, resp := retryingGet(t, client, "https://localhost:"+portOf(t, conf.Addr)+"/warmup") - _, _ = io.Copy(io.Discard, resp.Body) - _ = resp.Body.Close() - _, resp = retryingGet(t, client, "https://localhost:"+portOf(t, conf.Addr)+"/secure") - _, _ = io.Copy(io.Discard, resp.Body) - _ = resp.Body.Close() - - // TLS capture is ciphertext, so assert the run capture holds TLS - // records in both directions rather than plaintext content. - waitHTTPCapture(t, path, "") - waitTCPPayloads(t, path, func(joined string) bool { - return strings.Contains(joined, "\x16") && strings.Contains(joined, "\x17") - }) - - svc.Stop(context.Background()) //nolint:errcheck,gosec - discardStartErr(t, errCh) -} - -// retryingGet issues a GET with Connection: close (forcing the server to -// close the connection so its capture is flushed), retrying while the service -// is still starting up. -func retryingGet(t *testing.T, client *http.Client, url string) (status int, resp *http.Response) { - t.Helper() - deadline := time.Now().Add(3 * time.Second) - var lastErr error - for time.Now().Before(deadline) { - req, err := http.NewRequest(http.MethodGet, url, nil) - if err != nil { - t.Fatalf("NewRequest: %v", err) - } - req.Header.Set("Connection", "close") - r, err := client.Do(req) - if err == nil { - return r.StatusCode, r - } - lastErr = err - time.Sleep(20 * time.Millisecond) - } - t.Fatalf("GET %s: %v", url, lastErr) - return 0, nil + _ = svc.Stop(context.Background()) + testutil.DiscardServiceStartErr(t, errCh) } func startHTTPService(t *testing.T, conf Config, run *capture.Run) (service.Service, <-chan error) { @@ -149,94 +55,3 @@ func startHTTPService(t *testing.T, conf Config, run *capture.Run) (service.Serv go func() { errCh <- svc.Start(context.Background()) }() return svc, errCh } - -func discardStartErr(t *testing.T, errCh <-chan error) { - t.Helper() - select { - case err := <-errCh: - if err != nil { - t.Fatalf("service.Start returned error: %v", err) - } - case <-time.After(3 * time.Second): - t.Fatalf("service.Start never returned") - } -} - -// waitHTTPCapture waits for the run capture's payloads to contain want -// (or for any packets when want is empty). -func waitHTTPCapture(t *testing.T, path, want string) { - t.Helper() - deadline := time.Now().Add(3 * time.Second) - for time.Now().Before(deadline) { - if payloads, err := tcpPayloads(path); err == nil && (want == "" || strings.Contains(payloads, want)) { - return - } - time.Sleep(20 * time.Millisecond) - } - t.Fatalf("run capture %s never contained %q", path, want) -} - -func waitTCPPayloadsContain(t *testing.T, path, want string) { - t.Helper() - waitTCPPayloads(t, path, func(joined string) bool { - return strings.Contains(joined, want) - }) -} - -func waitTCPPayloads(t *testing.T, path string, cond func(string) bool) { - t.Helper() - deadline := time.Now().Add(3 * time.Second) - for time.Now().Before(deadline) { - joined, err := tcpPayloads(path) - if err == nil && cond(joined) { - return - } - time.Sleep(20 * time.Millisecond) - } - joined, _ := tcpPayloads(path) - t.Fatalf("capture %s never satisfied predicate (payloads %q)", path, joined) -} - -func tcpPayloads(path string) (string, error) { - f, err := os.Open(path) - if err != nil { - return "", err - } - defer f.Close() - r, err := pcapgo.NewNgReader(f, pcapgo.DefaultNgReaderOptions) - if err != nil { - return "", err - } - var sb strings.Builder - for { - data, _, err := r.ReadPacketData() - if err != nil { - break - } - pkt := gopacket.NewPacket(data, layers.LinkTypeEthernet, gopacket.Default) - if t, ok := pkt.Layer(layers.LayerTypeTCP).(*layers.TCP); ok { - sb.Write(t.Payload) - } - } - return sb.String(), nil -} - -func freeTCPAddr(t *testing.T) string { - t.Helper() - ln, err := net.Listen("tcp", "127.0.0.1:0") - if err != nil { - t.Fatalf("Listen: %v", err) - } - addr := ln.Addr().String() - _ = ln.Close() - return addr -} - -func portOf(t *testing.T, addr string) string { - t.Helper() - _, port, err := net.SplitHostPort(addr) - if err != nil { - t.Fatalf("SplitHostPort(%q): %v", addr, err) - } - return port -} diff --git a/internal/httpserver/fakemode.go b/internal/httpserver/fakemode.go index aa203da..1c7e9c1 100644 --- a/internal/httpserver/fakemode.go +++ b/internal/httpserver/fakemode.go @@ -50,10 +50,6 @@ type fakeResponse struct { type fakeGenerator func(r *http.Request, m fakeMeta) fakeResponse -// statusOverrideWriter forces a configured status code, but only for ordinary -// 200 OK responses. It deliberately leaves partial content (206) and -// not-modified (304) responses untouched so conditional/range requests still -// behave correctly in real mode. type statusOverrideWriter struct { http.ResponseWriter status int @@ -83,8 +79,6 @@ func (w *statusCaptureWriter) WriteHeader(code int) { } func (h FakeHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { - logger := h.Logger - m := resolveFakeMeta(r.URL.Path) gen := defaultFakeRegistry.lookup(m.ext) resp := gen(r, m) @@ -93,28 +87,7 @@ func (h FakeHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", resp.contentType) } - cap := &statusCaptureWriter{ResponseWriter: w} - out := http.ResponseWriter(cap) - if h.StatusCode != 0 { - out = &statusOverrideWriter{ResponseWriter: cap, status: h.StatusCode} - } - - http.ServeContent(out, r, m.name, resp.modTime, bytes.NewReader(resp.body)) - - status := cap.status - if status == 0 { - // ServeContent defaults to 200 if it wrote a body. - status = 200 - } - logger.Info( - r.Method, - "src", r.RemoteAddr, - "to", r.URL.Path, - "status", status, - "host", r.Host, - "ua", r.UserAgent(), - "len", r.ContentLength, - ) + serveContent(w, r, m.name, resp.modTime, bytes.NewReader(resp.body), h.StatusCode, h.Logger, r.ContentLength) } func resolveFakeMeta(urlPath string) fakeMeta { diff --git a/internal/httpserver/http_test.go b/internal/httpserver/http_test.go index 757964d..3f1fe44 100644 --- a/internal/httpserver/http_test.go +++ b/internal/httpserver/http_test.go @@ -162,28 +162,13 @@ func TestHTTPSServer_Smoke(t *testing.T) { func mustGet(t *testing.T, client *http.Client, url string) *http.Response { t.Helper() - - deadline := time.Now().Add(2 * time.Second) - var lastErr error - for time.Now().Before(deadline) { - resp, err := client.Get(url) - if err == nil { - return resp - } - lastErr = err - time.Sleep(10 * time.Millisecond) - } - t.Fatalf("GET %s: %v", url, lastErr) - return nil + _, resp := testutil.RetryGet(t, client, url) + return resp } func portFromAddr(t *testing.T, addr string) string { t.Helper() - _, port, err := net.SplitHostPort(addr) - if err != nil { - t.Fatalf("SplitHostPort(%q): %v", addr, err) - } - return port + return testutil.MustPort(t, addr) } // tempDirWithFiles creates a temporary directory, writes the given files into it, @@ -258,98 +243,6 @@ func TestRealHandler_ServesHTMLFile(t *testing.T) { } } -func TestRealHandler_ServesTextFile(t *testing.T) { - dir := tempDirWithFiles(t, map[string]string{ - "readme.txt": "this is a plain text file", - }) - _, base := startRealServer(t, dir, 0) - - resp := mustGet(t, http.DefaultClient, base+"/readme.txt") - defer resp.Body.Close() //nolint:errcheck - - if resp.StatusCode != http.StatusOK { - t.Fatalf("expected 200, got %d", resp.StatusCode) - } - if ct := resp.Header.Get("Content-Type"); !strings.HasPrefix(ct, "text/plain") { - t.Fatalf("expected text/plain Content-Type, got %q", ct) - } - body, _ := io.ReadAll(resp.Body) - if !strings.Contains(string(body), "plain text file") { - t.Fatalf("unexpected body: %q", string(body)) - } -} - -func TestRealHandler_ServesBinaryFile(t *testing.T) { - // A minimal valid PNG (1x1 pixel, transparent) - pngBytes := []byte{ - 0x89, 0x50, 0x4e, 0x47, 0x0d, 0x0a, 0x1a, 0x0a, - 0x00, 0x00, 0x00, 0x0d, 0x49, 0x48, 0x44, 0x52, - 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01, - 0x08, 0x06, 0x00, 0x00, 0x00, 0x1f, 0x15, 0xc4, - 0x89, 0x00, 0x00, 0x00, 0x0b, 0x49, 0x44, 0x41, - 0x54, 0x08, 0xd7, 0x63, 0x60, 0x00, 0x00, 0x00, - 0x02, 0x00, 0x01, 0xe2, 0x21, 0xbc, 0x33, 0x00, - 0x00, 0x00, 0x00, 0x49, 0x45, 0x4e, 0x44, 0xae, - 0x42, 0x60, 0x82, - } - dir := t.TempDir() - if err := os.WriteFile(filepath.Join(dir, "pixel.png"), pngBytes, 0o644); err != nil { - t.Fatalf("WriteFile: %v", err) - } - _, base := startRealServer(t, dir, 0) - - resp := mustGet(t, http.DefaultClient, base+"/pixel.png") - defer resp.Body.Close() //nolint:errcheck - - if resp.StatusCode != http.StatusOK { - t.Fatalf("expected 200, got %d", resp.StatusCode) - } - if ct := resp.Header.Get("Content-Type"); !strings.HasPrefix(ct, "image/png") { - t.Fatalf("expected image/png Content-Type, got %q", ct) - } - body, _ := io.ReadAll(resp.Body) - if len(body) != len(pngBytes) { - t.Fatalf("expected %d bytes, got %d", len(pngBytes), len(body)) - } -} - -func TestRealHandler_ServesFileFromSubdirectory(t *testing.T) { - dir := tempDirWithFiles(t, map[string]string{ - "assets/style.css": "body { color: red; }", - }) - _, base := startRealServer(t, dir, 0) - - resp := mustGet(t, http.DefaultClient, base+"/assets/style.css") - defer resp.Body.Close() //nolint:errcheck - - if resp.StatusCode != http.StatusOK { - t.Fatalf("expected 200, got %d", resp.StatusCode) - } - body, _ := io.ReadAll(resp.Body) - if !strings.Contains(string(body), "color: red") { - t.Fatalf("unexpected body: %q", string(body)) - } -} - -func TestRealHandler_RootPathServesIndexHTML(t *testing.T) { - dir := tempDirWithFiles(t, map[string]string{ - "index.html": "root index", - }) - _, base := startRealServer(t, dir, 0) - - resp := mustGet(t, http.DefaultClient, base+"/") - defer resp.Body.Close() //nolint:errcheck - - // / should fall through to index.html - if resp.StatusCode != http.StatusOK { - t.Fatalf("expected 200, got %d", resp.StatusCode) - } - body, _ := io.ReadAll(resp.Body) - if !strings.Contains(string(body), "root index") { - t.Fatalf("unexpected body: %q", string(body)) - } -} - func TestRealHandler_StatusCodeOverride(t *testing.T) { dir := tempDirWithFiles(t, map[string]string{ "page.html": "ok", @@ -378,110 +271,29 @@ func TestRealHandler_MissingFileReturns404(t *testing.T) { } } -func TestRealHandler_DirectoryRequestWithoutIndexReturns404(t *testing.T) { - dir := tempDirWithFiles(t, map[string]string{ - "sub/file.txt": "content", - }) - _, base := startRealServer(t, dir, 0) - - // Request the subdirectory itself — without an index.html it should 404, - // not serve a listing. - resp := mustGet(t, http.DefaultClient, base+"/sub/") - defer resp.Body.Close() //nolint:errcheck - - if resp.StatusCode != http.StatusNotFound { - t.Fatalf("expected 404 for directory request, got %d", resp.StatusCode) - } -} - -func TestRealHandler_DirectoryServesIndexHTML(t *testing.T) { - dir := tempDirWithFiles(t, map[string]string{ - "sub/index.html": "sub index", - }) - _, base := startRealServer(t, dir, 0) - - resp := mustGet(t, http.DefaultClient, base+"/sub/") - defer resp.Body.Close() //nolint:errcheck - - if resp.StatusCode != http.StatusOK { - t.Fatalf("expected 200, got %d", resp.StatusCode) - } - body, _ := io.ReadAll(resp.Body) - if !strings.Contains(string(body), "sub index") { - t.Fatalf("unexpected body: %q", string(body)) - } -} - -func TestRealHandler_ConditionalRequestNotOverridden(t *testing.T) { - dir := tempDirWithFiles(t, map[string]string{ - "page.txt": "hello", - }) - _, base := startRealServer(t, dir, http.StatusAccepted) // non-200 override - - // Request once to learn the Last-Modified, then re-request with - // If-Modified-Since to trigger a 304. The configured status override must - // not clobber the 304 Not Modified response. - first := mustGet(t, http.DefaultClient, base+"/page.txt") - lm := first.Header.Get("Last-Modified") - _ = first.Body.Close() - if lm == "" { - t.Fatal("expected a Last-Modified header on the first response") - } - - req, err := http.NewRequest(http.MethodGet, base+"/page.txt", nil) - if err != nil { - t.Fatal(err) - } - req.Header.Set("If-Modified-Since", lm) - resp, err := http.DefaultClient.Do(req) - if err != nil { - t.Fatal(err) - } - defer resp.Body.Close() //nolint:errcheck - - if resp.StatusCode != http.StatusNotModified { - t.Fatalf("expected 304 for conditional request, got %d", resp.StatusCode) - } -} - // --- security tests --- func TestRealHandler_TraversalBlocked(t *testing.T) { - paths := []struct { - name string - path string - }{ - {"classic", "/../secret.txt"}, - {"encoded", "/%2e%2e/secret.txt"}, - {"middle", "/a/../secret.txt"}, - {"raw", "/../../secret.txt"}, + // Single server; sentinel file lives outside the root. + parent := t.TempDir() + secret := filepath.Join(parent, "secret.txt") + if err := os.WriteFile(secret, []byte("secret contents"), 0o644); err != nil { + t.Fatalf("WriteFile: %v", err) } - for _, tc := range paths { - t.Run(tc.name, func(t *testing.T) { - // Write a sentinel file one level above the root dir. - parent := t.TempDir() - secret := filepath.Join(parent, "secret.txt") - if err := os.WriteFile(secret, []byte("secret contents"), 0o644); err != nil { - t.Fatalf("WriteFile: %v", err) - } - - // The server root is a subdirectory; secret.txt is outside it. - root := filepath.Join(parent, "www") - if err := os.MkdirAll(root, 0o755); err != nil { - t.Fatalf("MkdirAll: %v", err) - } - - _, base := startRealServer(t, root, 0) - - resp := mustGet(t, http.DefaultClient, base+tc.path) - defer resp.Body.Close() //nolint:errcheck - - // Must not serve the file — 404 or 400 are both acceptable - if resp.StatusCode == http.StatusOK { - body, _ := io.ReadAll(resp.Body) - t.Fatalf("traversal %q succeeded — got 200 with body: %q", tc.path, string(body)) - } - }) + root := filepath.Join(parent, "www") + if err := os.MkdirAll(root, 0o755); err != nil { + t.Fatalf("MkdirAll: %v", err) + } + _, base := startRealServer(t, root, 0) + + for _, path := range []string{"/../secret.txt", "/%2e%2e/secret.txt", "/a/../secret.txt", "/../../secret.txt"} { + resp := mustGet(t, http.DefaultClient, base+path) + body, _ := io.ReadAll(resp.Body) + _ = resp.Body.Close() + // Mmst not serve the file, 404 or 400 are both accepted + if resp.StatusCode == http.StatusOK { + t.Fatalf("traversal %q succeeded — got 200 with body: %q", path, string(body)) + } } } @@ -498,15 +310,3 @@ func TestNewServer_RealMode_MissingRootDirReturnsError(t *testing.T) { t.Fatal("expected error when RootDir is empty, got nil") } } - -func TestNewServer_RealMode_NonexistentRootDirReturnsError(t *testing.T) { - logger := testutil.Logger() - _, err := NewServer(Config{ - Addr: "127.0.0.1:0", - Mode: "real", - RootDir: "/this/path/does/not/exist", - }, nil, logger) - if err == nil { - t.Fatal("expected error for nonexistent RootDir, got nil") - } -} diff --git a/internal/httpserver/realmode.go b/internal/httpserver/realmode.go index 115099b..d4f3e33 100644 --- a/internal/httpserver/realmode.go +++ b/internal/httpserver/realmode.go @@ -1,7 +1,6 @@ package httpserver import ( - "fmt" "log/slog" "net/http" "os" @@ -97,27 +96,7 @@ func (h RealHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { } defer f.Close() //nolint:errcheck - cap := &statusCaptureWriter{ResponseWriter: w} - out := http.ResponseWriter(cap) - if h.StatusCode != 0 { - out = &statusOverrideWriter{ResponseWriter: cap, status: h.StatusCode} - } - - http.ServeContent(out, r, stat.Name(), stat.ModTime(), f) - - status := cap.status - if status == 0 { - status = http.StatusOK - } - logger.Info( - r.Method, - "src", r.RemoteAddr, - "to", r.URL.Path, - "status", status, - "host", r.Host, - "ua", r.UserAgent(), - "len", fmt.Sprintf("%d", stat.Size()), - ) + serveContent(w, r, stat.Name(), stat.ModTime(), f, h.StatusCode, logger, stat.Size()) } // pathWithin reports whether child is inside parent (or equals it), using only diff --git a/internal/httpserver/server.go b/internal/httpserver/server.go index 19f8ff8..4c47729 100644 --- a/internal/httpserver/server.go +++ b/internal/httpserver/server.go @@ -4,6 +4,7 @@ import ( "context" "crypto/tls" "errors" + "io" "log/slog" "net/http" "strings" @@ -101,3 +102,27 @@ func (s *Server) Stop(ctx context.Context) error { } return nil } + +func serveContent(w http.ResponseWriter, r *http.Request, name string, modTime time.Time, content io.ReadSeeker, statusOverride int, logger *slog.Logger, contentLen any) { + cap := &statusCaptureWriter{ResponseWriter: w} + out := http.ResponseWriter(cap) + if statusOverride != 0 { + out = &statusOverrideWriter{ResponseWriter: cap, status: statusOverride} + } + + http.ServeContent(out, r, name, modTime, content) + + status := cap.status + if status == 0 { + status = http.StatusOK + } + logger.Info( + r.Method, + "src", r.RemoteAddr, + "to", r.URL.Path, + "status", status, + "host", r.Host, + "ua", r.UserAgent(), + "len", contentLen, + ) +} diff --git a/internal/listener/listener_test.go b/internal/listener/listener_test.go index 9022425..01ad342 100644 --- a/internal/listener/listener_test.go +++ b/internal/listener/listener_test.go @@ -10,12 +10,10 @@ package listener import ( "context" - "crypto/tls" "io" "log/slog" "net" "os" - "path/filepath" "strings" "testing" "time" @@ -23,7 +21,6 @@ import ( "github.com/google/gopacket" "github.com/google/gopacket/layers" "github.com/google/gopacket/pcapgo" - "github.com/lachlanharrisdev/gonetsim/internal/capture" "github.com/lachlanharrisdev/gonetsim/internal/service" "github.com/lachlanharrisdev/gonetsim/internal/testutil" "github.com/lachlanharrisdev/gonetsim/internal/tlsprovider" @@ -33,22 +30,6 @@ func testLogger() *slog.Logger { return testutil.Logger() } -func testRun(t *testing.T) (*capture.Run, string) { - t.Helper() - path := filepath.Join(t.TempDir(), "run.pcapng") - run, err := capture.NewRun(path) - if err != nil { - t.Fatalf("NewRun: %v", err) - } - t.Cleanup(func() { _ = run.Close() }) - return run, path -} - -func freePort(t *testing.T, network string) string { - t.Helper() - return testutil.FreePort(t, network) -} - func startService(t *testing.T, svc service.Service) { t.Helper() ctx, cancel := context.WithCancel(context.Background()) @@ -85,28 +66,11 @@ func dialTCP(t *testing.T, addr string) net.Conn { return nil } -func dialTLS(t *testing.T, addr, serverName string) net.Conn { - t.Helper() - tlsConf := &tls.Config{InsecureSkipVerify: true, ServerName: serverName} - deadline := time.Now().Add(2 * time.Second) - var conn net.Conn - var lastErr error - for time.Now().Before(deadline) { - conn, lastErr = tls.Dial("tcp", addr, tlsConf) - if lastErr == nil { - return conn - } - time.Sleep(10 * time.Millisecond) - } - t.Fatalf("tls.Dial %s failed: %v", addr, lastErr) - return nil -} - func echoConfig(t *testing.T) Config { return Config{ Name: "echotest", Network: "tcp", - Addr: freePort(t, "tcp"), + Addr: testutil.FreePort(t, "tcp"), HandlerSpec: "builtin:echo", ReadTimeout: 5 * time.Second, Capture: true, @@ -117,32 +81,20 @@ func echoConfig(t *testing.T) Config { // matches want, tolerating the async flush that follows connection teardown. func waitTransportFrames(t *testing.T, path, proto string, want []string) { t.Helper() - deadline := time.Now().Add(3 * time.Second) - for time.Now().Before(deadline) { + testutil.WaitFor(t, 3*time.Second, "payload sequence match", func() bool { got, err := transportPayloads(path, proto) - if err == nil && strings.Join(got, "|") == strings.Join(want, "|") { - return - } - time.Sleep(20 * time.Millisecond) - } - got, _ := transportPayloads(path, proto) - t.Fatalf("payload sequence never matched %q (proto %s), last saw %q", strings.Join(want, "|"), proto, strings.Join(got, "|")) + return err == nil && strings.Join(got, "|") == strings.Join(want, "|") + }) } // waitSubstringFrames polls until the concatenated payload sequence of a // capture contains want (used where multiple datagrams share one writer). func waitSubstringFrames(t *testing.T, path, proto, want string) { t.Helper() - deadline := time.Now().Add(3 * time.Second) - for time.Now().Before(deadline) { + testutil.WaitFor(t, 3*time.Second, "payload substring match", func() bool { got, err := transportPayloads(path, proto) - if err == nil && strings.Contains(strings.Join(got, "|"), want) { - return - } - time.Sleep(20 * time.Millisecond) - } - got, _ := transportPayloads(path, proto) - t.Fatalf("payload sequence never contained %q (proto %s), last saw %q", want, proto, strings.Join(got, "|")) + return err == nil && strings.Contains(strings.Join(got, "|"), want) + }) } // transportPayloads extracts transport-layer payloads from a pcapng file, @@ -180,7 +132,7 @@ func transportPayloads(path, proto string) ([]string, error) { func TestTCPService(t *testing.T) { t.Run("echo over tcp with pcapng capture", func(t *testing.T) { conf := echoConfig(t) - run, path := testRun(t) + run, path := testutil.NewPcapRun(t) svc, err := NewService(conf, nil, testLogger(), run) if err != nil { @@ -201,10 +153,8 @@ func TestTCPService(t *testing.T) { } _ = conn.Close() - // produce a pcapng file with the exchanged payload - if _, err := transportPayloads(path, "tcp"); err != nil { - t.Fatalf("capture: %v", err) - } + waitTransportFrames(t, path, "tcp", + []string{"", "", "hello\n", "hello\n", "", ""}) }) t.Run("idle timeout closes connection", func(t *testing.T) { @@ -231,7 +181,7 @@ func TestTCPService(t *testing.T) { conf := Config{ Name: "isotest", Network: "tcp", - Addr: freePort(t, "tcp"), + Addr: testutil.FreePort(t, "tcp"), HandlerSpec: "lua:isolated.lua", BaseDir: "../handler/testdata", ReadTimeout: 5 * time.Second, @@ -262,58 +212,6 @@ func TestTCPService(t *testing.T) { }) } -func TestTCPCapture(t *testing.T) { - t.Run("tcp listener produces pcapng with frames", func(t *testing.T) { - conf := echoConfig(t) - run, path := testRun(t) - - svc, err := NewService(conf, nil, testLogger(), run) - if err != nil { - t.Fatalf("NewService: %v", err) - } - startService(t, svc) - - conn := dialTCP(t, conf.Addr) - if _, err := conn.Write([]byte("hello\n")); err != nil { - t.Fatalf("Write: %v", err) - } - buf := make([]byte, 6) - if _, err := io.ReadFull(conn, buf); err != nil { - t.Fatalf("ReadFull: %v", err) - } - _ = conn.Close() - - waitTransportFrames(t, path, "tcp", - []string{"", "", "hello\n", "hello\n", "", ""}) - }) -} - -func TestTCPServiceTLS(t *testing.T) { - t.Run("echo over TLS", func(t *testing.T) { - conf := echoConfig(t) - conf.TLS = &tlsprovider.Config{} - - svc, err := NewService(conf, nil, testLogger(), nil) - if err != nil { - t.Fatalf("NewService: %v", err) - } - startService(t, svc) - - conn := dialTLS(t, conf.Addr, "localhost") - defer func() { _ = conn.Close() }() - if _, err := conn.Write([]byte("secure")); err != nil { - t.Fatalf("Write: %v", err) - } - buf := make([]byte, 6) - if _, err := io.ReadFull(conn, buf); err != nil { - t.Fatalf("ReadFull: %v", err) - } - if string(buf) != "secure" { - t.Fatalf("expected echo, got %q", buf) - } - }) -} - func TestUDPCapture(t *testing.T) { exchange := func(t *testing.T, addr, payload, want string) { t.Helper() @@ -349,12 +247,12 @@ func TestUDPCapture(t *testing.T) { conf := Config{ Name: "udpecho", Network: "udp", - Addr: freePort(t, "udp"), + Addr: testutil.FreePort(t, "udp"), HandlerSpec: "builtin:echo", ReadTimeout: 150 * time.Millisecond, Capture: true, } - run, path := testRun(t) + run, path := testutil.NewPcapRun(t) svc, err := NewService(conf, nil, testLogger(), run) if err != nil { t.Fatalf("NewService: %v", err) @@ -369,7 +267,7 @@ func TestUDPCapture(t *testing.T) { conf := Config{ Name: "udplua", Network: "udp", - Addr: freePort(t, "udp"), + Addr: testutil.FreePort(t, "udp"), HandlerSpec: "lua:packet.lua", BaseDir: "../handler/testdata", ReadTimeout: 5 * time.Second, @@ -388,7 +286,7 @@ func TestStartWithCancelledContext(t *testing.T) { conf := Config{ Name: "canceled-" + network, Network: network, - Addr: freePort(t, network), + Addr: testutil.FreePort(t, network), HandlerSpec: "builtin:sink", ReadTimeout: 5 * time.Second, } diff --git a/internal/listener/service.go b/internal/listener/service.go index a721950..0369ca3 100644 --- a/internal/listener/service.go +++ b/internal/listener/service.go @@ -29,5 +29,5 @@ func NewService(conf Config, global *state.Store, logger *slog.Logger, run *capt if conf.Network == "udp" { return &udpService{conf: conf, handler: h, log: log, run: run, global: global}, nil } - return &tcpService{conf: conf, handler: h, log: log, run: run, global: global, idle: conf.ReadTimeout}, nil + return &tcpService{conf: conf, handler: h, log: log, run: run, global: global}, nil } diff --git a/internal/listener/tcp.go b/internal/listener/tcp.go index a357e48..274592a 100644 --- a/internal/listener/tcp.go +++ b/internal/listener/tcp.go @@ -22,7 +22,6 @@ type tcpService struct { log *slog.Logger run *capture.Run global *state.Store - idle time.Duration mu sync.Mutex ln net.Listener diff --git a/internal/service/manager.go b/internal/service/manager.go index ea85121..2a53fc6 100644 --- a/internal/service/manager.go +++ b/internal/service/manager.go @@ -32,10 +32,6 @@ func (m *Manager) RunAll(ctx context.Context) error { return runServices(ctx, m.logger, m.shutdownTimeout, m.services) } -func (m *Manager) RunSingleService(ctx context.Context, s Service) error { - return runServices(ctx, m.logger, m.shutdownTimeout, []Service{s}) -} - func runServices(ctx context.Context, logger *slog.Logger, shutdownTimeout time.Duration, services []Service) error { if len(services) == 0 { return nil diff --git a/internal/service/manager_test.go b/internal/service/manager_test.go deleted file mode 100644 index 0a7cc3e..0000000 --- a/internal/service/manager_test.go +++ /dev/null @@ -1,98 +0,0 @@ -////---------------------------------------------------------------------------- -// NOTICE: to save development time, test files (including this) have been -// generated with LLMs. The author(s) do not claim credit for these tests -// and exist purely for maximising code quality and reliability -// -// For more information please see `/.github/AI_USAGE.md` -//----------------------------------------------------------------------------// - -package service - -import ( - "context" - "errors" - "log/slog" - "strings" - "testing" - "time" - - "github.com/lachlanharrisdev/gonetsim/internal/testutil" -) - -type fakeService struct { - name string - startErr error - started chan struct{} - stopped chan struct{} - block bool -} - -func (f *fakeService) Name() string { return f.name } - -func (f *fakeService) Start(ctx context.Context) error { - close(f.started) - if f.block { - <-ctx.Done() - } - return f.startErr -} - -func (f *fakeService) Stop(ctx context.Context) error { - close(f.stopped) - return nil -} - -func discardLogger() *slog.Logger { - return testutil.Logger() -} - -func TestRunServices_PropagatesStartError(t *testing.T) { - svc := &fakeService{ - name: "boom", - startErr: errors.New("bind: address already in use"), - started: make(chan struct{}), - stopped: make(chan struct{}), - } - - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - - err := runServices(ctx, discardLogger(), time.Second, []Service{svc}) - if err == nil { - t.Fatal("expected an error, got nil") - } - if !strings.Contains(err.Error(), "boom") { - t.Fatalf("expected error to name the failing service, got %q", err) - } - if !strings.Contains(err.Error(), "address already in use") { - t.Fatalf("expected the underlying error to be preserved, got %q", err) - } -} - -func TestRunServices_ReturnsNilOnCleanShutdown(t *testing.T) { - svc := &fakeService{ - name: "ok", - block: true, - started: make(chan struct{}), - stopped: make(chan struct{}), - } - - ctx, cancel := context.WithCancel(context.Background()) - - done := make(chan error, 1) - go func() { - done <- runServices(ctx, discardLogger(), time.Second, []Service{svc}) - }() - - <-svc.started - cancel() - - select { - case err := <-done: - if err != nil { - t.Fatalf("expected nil on clean shutdown, got %v", err) - } - case <-time.After(5 * time.Second): - t.Fatal("manager did not return after cancellation") - } -} diff --git a/internal/state/state.go b/internal/state/state.go index fefb5dd..18af17b 100644 --- a/internal/state/state.go +++ b/internal/state/state.go @@ -48,9 +48,7 @@ func (s *Store) Get(key string) (string, bool) { } func (s *Store) Has(key string) bool { - s.budget.mu.RLock() - defer s.budget.mu.RUnlock() - _, ok := s.data[key] + _, ok := s.Get(key) return ok } diff --git a/internal/testutil/testutil.go b/internal/testutil/testutil.go index 8ab96aa..d13cc75 100644 --- a/internal/testutil/testutil.go +++ b/internal/testutil/testutil.go @@ -4,7 +4,19 @@ import ( "io" "log/slog" "net" + "net/http" + "os" + "path/filepath" + "strings" "testing" + "time" + + "github.com/google/gopacket" + "github.com/google/gopacket/layers" + "github.com/google/gopacket/pcapgo" + "github.com/miekg/dns" + + "github.com/lachlanharrisdev/gonetsim/internal/capture" ) func Logger() *slog.Logger { @@ -33,3 +45,131 @@ func FreePort(t *testing.T, network string) string { } return FreeTCPAddr(t) } + +func MustPort(t *testing.T, addr string) string { + t.Helper() + _, port, err := net.SplitHostPort(addr) + if err != nil { + t.Fatalf("SplitHostPort(%q): %v", addr, err) + } + return port +} + +func NewPcapRun(t *testing.T) (*capture.Run, string) { + t.Helper() + path := filepath.Join(t.TempDir(), "run.pcapng") + run, err := capture.NewRun(path) + if err != nil { + t.Fatalf("NewRun: %v", err) + } + t.Cleanup(func() { _ = run.Close() }) + return run, path +} + +func WaitFor(t *testing.T, timeout time.Duration, msg string, cond func() bool) { + t.Helper() + deadline := time.Now().Add(timeout) + for time.Now().Before(deadline) { + if cond() { + return + } + time.Sleep(20 * time.Millisecond) + } + t.Fatalf("timed out waiting: %s", msg) +} + +func TransportPayloads(path string) (string, error) { + f, err := os.Open(path) + if err != nil { + return "", err + } + defer f.Close() + r, err := pcapgo.NewNgReader(f, pcapgo.DefaultNgReaderOptions) + if err != nil { + return "", err + } + var sb strings.Builder + for { + data, _, err := r.ReadPacketData() + if err != nil { + break + } + pkt := gopacket.NewPacket(data, layers.LinkTypeEthernet, gopacket.Default) + if u, ok := pkt.Layer(layers.LayerTypeUDP).(*layers.UDP); ok { + sb.Write(u.Payload) + } else if tc, ok := pkt.Layer(layers.LayerTypeTCP).(*layers.TCP); ok { + sb.Write(tc.Payload) + } + } + return sb.String(), nil +} + +func WaitForPayload(t *testing.T, path string, timeout time.Duration, cond func(string) bool) string { + t.Helper() + deadline := time.Now().Add(timeout) + for time.Now().Before(deadline) { + if joined, err := TransportPayloads(path); err == nil && cond(joined) { + return joined + } + time.Sleep(20 * time.Millisecond) + } + joined, _ := TransportPayloads(path) + t.Fatalf("capture %s never satisfied predicate (payloads %q)", path, joined) + return "" +} + +func WaitForPayloadContains(t *testing.T, path, want string, timeout time.Duration) { + t.Helper() + WaitForPayload(t, path, timeout, func(s string) bool { + return want == "" || strings.Contains(s, want) + }) +} + +func DiscardServiceStartErr(t *testing.T, errCh <-chan error) { + t.Helper() + select { + case err := <-errCh: + if err != nil { + t.Fatalf("service.Start returned error: %v", err) + } + case <-time.After(3 * time.Second): + t.Fatalf("service.Start never returned") + } +} + +func RetryGet(t *testing.T, client *http.Client, url string) (int, *http.Response) { + t.Helper() + deadline := time.Now().Add(3 * time.Second) + var lastErr error + for time.Now().Before(deadline) { + req, err := http.NewRequest(http.MethodGet, url, nil) + if err != nil { + t.Fatalf("NewRequest: %v", err) + } + req.Header.Set("Connection", "close") + r, err := client.Do(req) + if err == nil { + return r.StatusCode, r + } + lastErr = err + time.Sleep(20 * time.Millisecond) + } + t.Fatalf("GET %s: %v", url, lastErr) + return 0, nil +} + +func RetryDNSExchange(t *testing.T, client *dns.Client, addr string, m *dns.Msg) (*dns.Msg, time.Duration, error) { + t.Helper() + deadline := time.Now().Add(3 * time.Second) + var lastErr error + var lastRTT time.Duration + for time.Now().Before(deadline) { + resp, rtt, err := client.Exchange(m, addr) + if err == nil && resp != nil { + return resp, rtt, nil + } + lastErr, lastRTT = err, rtt + time.Sleep(20 * time.Millisecond) + } + return nil, lastRTT, lastErr +} diff --git a/internal/tlsprovider/tls_test.go b/internal/tlsprovider/tls_test.go index 4aae8f2..c32031b 100644 --- a/internal/tlsprovider/tls_test.go +++ b/internal/tlsprovider/tls_test.go @@ -28,35 +28,28 @@ func TestGenerateSelfSigned_SaneCertificate(t *testing.T) { ValidFor: 2 * time.Hour, }) if err != nil { - // failed with error t.Fatalf("GenerateSelfSigned: %v", err) } if len(cert.Certificate) == 0 { - // failed to generate certificate t.Fatalf("expected at least one certificate") } if cert.PrivateKey == nil { - // failed to generate private key t.Fatalf("expected PrivateKey to be set") } leaf, err := x509.ParseCertificate(cert.Certificate[0]) if err != nil { - // failed to parse generated certificate with error t.Fatalf("ParseCertificate: %v", err) } if time.Until(leaf.NotAfter) <= 0 { - // failed to generate a certificate that is currently valid t.Fatalf("expected certificate to be currently valid") } if leaf.KeyUsage&(x509.KeyUsageDigitalSignature|x509.KeyUsageKeyEncipherment) == 0 { - // failed to generate a certificate with appropriate key usage for TLS server t.Fatalf("expected KeyUsage to include digital signature and/or key encipherment, got %v", leaf.KeyUsage) } if len(leaf.ExtKeyUsage) == 0 || leaf.ExtKeyUsage[0] != x509.ExtKeyUsageServerAuth { - // failed to generate a certificate with appropriate extended key usage for TLS server t.Fatalf("expected ExtKeyUsage to include server auth, got %v", leaf.ExtKeyUsage) } @@ -77,7 +70,7 @@ func TestGenerateSelfSigned_SaneCertificate(t *testing.T) { } -func TestTLSConfig_AutoPersistedPair_Reused(t *testing.T) { +func TestTLSConfig_PersistReuseRegenerate(t *testing.T) { dir := t.TempDir() cfg := Config{ @@ -103,6 +96,7 @@ func TestTLSConfig_AutoPersistedPair_Reused(t *testing.T) { t.Fatalf("ReadFile(ca): %v", err) } + // Second load must reuse the persisted pair. _, err = cfg.TLSConfig() if err != nil { t.Fatalf("TLSConfig (second): %v", err) @@ -130,24 +124,8 @@ func TestTLSConfig_AutoPersistedPair_Reused(t *testing.T) { if !bytes.Equal(ca1, ca2) { t.Fatalf("expected CA to be reused") } -} - -func TestTLSConfig_RegenerateForce(t *testing.T) { - dir := t.TempDir() - - cfg := Config{ - CertFile: filepath.Join(dir, PersistedCertFileName), - KeyFile: filepath.Join(dir, PersistedKeyFileName), - } - - if _, err := cfg.TLSConfig(); err != nil { - t.Fatalf("TLSConfig (initial): %v", err) - } - before, err := os.ReadFile(cfg.CertFile) - if err != nil { - t.Fatalf("ReadFile(cert): %v", err) - } + // Force regeneration must produce a different cert. if err := cfg.Regenerate(); err != nil { t.Fatalf("Regenerate: %v", err) } @@ -158,18 +136,16 @@ func TestTLSConfig_RegenerateForce(t *testing.T) { if err != nil { t.Fatalf("ReadFile(cert, after): %v", err) } - - if bytes.Equal(before, after) { + if bytes.Equal(cert1, after) { t.Fatalf("expected cert to be regenerated, but it is identical") } -} -func TestCertExpired(t *testing.T) { - cert, err := GenerateSelfSigned(SelfSignedOptions{}) + // Freshly generated certs must not be expired. + fresh, err := GenerateSelfSigned(SelfSignedOptions{}) if err != nil { t.Fatalf("GenerateSelfSigned: %v", err) } - if certExpired(cert) { + if certExpired(fresh) { t.Fatalf("freshly generated cert must not be expired") } } From 7dd9be7685750bb84a759462e176714af24523d4 Mon Sep 17 00:00:00 2001 From: Lachlan Harris Date: Sun, 6 Sep 2026 18:18:50 +1000 Subject: [PATCH 4/4] fix: lint errors --- cmd/pcap.go | 16 ++++++++++++---- internal/capture/inspect.go | 2 +- internal/capture/run.go | 17 +++++++++-------- internal/capture/session_test.go | 2 +- internal/listener/listener_test.go | 2 +- internal/testutil/testutil.go | 2 +- 6 files changed, 25 insertions(+), 16 deletions(-) diff --git a/cmd/pcap.go b/cmd/pcap.go index face9f5..b7fce3f 100644 --- a/cmd/pcap.go +++ b/cmd/pcap.go @@ -45,7 +45,9 @@ func inspectPcap(out io.Writer, target string) error { if err != nil { return err } - fmt.Fprintf(out, "%s\n", summarizePcap(target, info)) + if _, err := fmt.Fprintf(out, "%s\n", summarizePcap(target, info)); err != nil { + return fmt.Errorf("write output: %w", err) + } return nil } @@ -89,14 +91,20 @@ func inspectPcapDir(out io.Writer, dir string) error { for _, f := range files { info, err := capture.Inspect(f) if err != nil { - fmt.Fprintf(out, "%s: ERROR %v\n", f, err) + if _, werr := fmt.Fprintf(out, "%s: ERROR %v\n", f, err); werr != nil { + return fmt.Errorf("write output: %w", werr) + } failed++ continue } - fmt.Fprintf(out, "%s\n", summarizePcap(f, info)) + if _, err := fmt.Fprintf(out, "%s\n", summarizePcap(f, info)); err != nil { + return fmt.Errorf("write output: %w", err) + } total += info.Packets } - fmt.Fprintf(out, "total: files=%d packets=%d\n", len(files), total) + if _, err := fmt.Fprintf(out, "total: files=%d packets=%d\n", len(files), total); err != nil { + return fmt.Errorf("write output: %w", err) + } if failed > 0 { return fmt.Errorf("%d of %d files could not be read", failed, len(files)) } diff --git a/internal/capture/inspect.go b/internal/capture/inspect.go index 0eda1a2..a2d55d3 100644 --- a/internal/capture/inspect.go +++ b/internal/capture/inspect.go @@ -50,7 +50,7 @@ func Inspect(path string) (FileInfo, error) { if err != nil { return FileInfo{}, fmt.Errorf("open %q: %w", path, err) } - defer f.Close() + defer func() { _ = f.Close() }() var magic [4]byte if _, err := io.ReadFull(f, magic[:]); err != nil { diff --git a/internal/capture/run.go b/internal/capture/run.go index 9e4d673..036502a 100644 --- a/internal/capture/run.go +++ b/internal/capture/run.go @@ -31,6 +31,11 @@ func NewRunID() string { } func DefaultRunsDir() (string, error) { + // XDG_DATA_HOME is honored on all platforms so tests can redirect the + // runs directory via t.Setenv (os.UserCacheDir ignores it on Windows). + if xdg := os.Getenv("XDG_DATA_HOME"); xdg != "" { + return filepath.Join(xdg, "gonetsim", "runs"), nil + } var base string switch runtime.GOOS { case "windows": @@ -46,15 +51,11 @@ func DefaultRunsDir() (string, error) { } base = filepath.Join(home, "Library", "Application Support") default: - if xdg := os.Getenv("XDG_DATA_HOME"); xdg != "" { - base = xdg - } else { - home, err := os.UserHomeDir() - if err != nil { - return "", err - } - base = filepath.Join(home, ".local", "share") + home, err := os.UserHomeDir() + if err != nil { + return "", err } + base = filepath.Join(home, ".local", "share") } return filepath.Join(base, "gonetsim", "runs"), nil } diff --git a/internal/capture/session_test.go b/internal/capture/session_test.go index e84e6bd..7c1e212 100644 --- a/internal/capture/session_test.go +++ b/internal/capture/session_test.go @@ -59,7 +59,7 @@ func readFrames(t *testing.T, path string) []gopacket.Packet { if err != nil { t.Fatalf("Open: %v", err) } - defer f.Close() + defer func() { _ = f.Close() }() r, err := pcapgo.NewNgReader(f, pcapgo.DefaultNgReaderOptions) if err != nil { t.Fatalf("NewNgReader: %v", err) diff --git a/internal/listener/listener_test.go b/internal/listener/listener_test.go index 01ad342..13501ab 100644 --- a/internal/listener/listener_test.go +++ b/internal/listener/listener_test.go @@ -104,7 +104,7 @@ func transportPayloads(path, proto string) ([]string, error) { if err != nil { return nil, err } - defer f.Close() + defer func() { _ = f.Close() }() r, err := pcapgo.NewNgReader(f, pcapgo.DefaultNgReaderOptions) if err != nil { return nil, err diff --git a/internal/testutil/testutil.go b/internal/testutil/testutil.go index d13cc75..54551ef 100644 --- a/internal/testutil/testutil.go +++ b/internal/testutil/testutil.go @@ -83,7 +83,7 @@ func TransportPayloads(path string) (string, error) { if err != nil { return "", err } - defer f.Close() + defer func() { _ = f.Close() }() r, err := pcapgo.NewNgReader(f, pcapgo.DefaultNgReaderOptions) if err != nil { return "", err