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/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/pcap.go b/cmd/pcap.go
new file mode 100644
index 0000000..b7fce3f
--- /dev/null
+++ b/cmd/pcap.go
@@ -0,0 +1,112 @@
+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: "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
+ }
+ if _, err := fmt.Fprintf(out, "%s\n", summarizePcap(target, info)); err != nil {
+ return fmt.Errorf("write output: %w", err)
+ }
+ 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 {
+ if _, werr := fmt.Fprintf(out, "%s: ERROR %v\n", f, err); werr != nil {
+ return fmt.Errorf("write output: %w", werr)
+ }
+ failed++
+ continue
+ }
+ if _, err := fmt.Fprintf(out, "%s\n", summarizePcap(f, info)); err != nil {
+ return fmt.Errorf("write output: %w", err)
+ }
+ total += info.Packets
+ }
+ 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))
+ }
+ 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..69e2793 100644
--- a/cmd/targets_test.go
+++ b/cmd/targets_test.go
@@ -1,17 +1,20 @@
package cmd
import (
- "io"
"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"
)
func testLogger() *slog.Logger {
- return slog.New(slog.NewTextHandler(io.Discard, nil))
+ return testutil.Logger()
}
func disabledAll(cfg *appconfig.Config) {
@@ -104,10 +107,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 +122,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 +144,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 +152,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 +161,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 +196,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 +233,35 @@ 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
+}
+
+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 2cb89b4..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
-
- # listener capture files are written to ./artifacts inside the container
- # (i.e. /artifacts); mount a volume there to keep them:
- # volumes:
- # - ./artifacts:/artifacts
+ # - ./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/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 2e002ae..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
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/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..a2d55d3
--- /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 func() { _ = 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/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..036502a
--- /dev/null
+++ b/internal/capture/run.go
@@ -0,0 +1,189 @@
+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) {
+ // 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":
+ 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:
+ 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..314317b
--- /dev/null
+++ b/internal/capture/session.go
@@ -0,0 +1,301 @@
+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
+
+ // 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 // comment attached to the next emitted frame
+
+ 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 = ipLayer(src, dst, layers.IPProtocolTCP, tcp)
+ transport = tcp
+ default:
+ udp := &layers.UDP{SrcPort: layers.UDPPort(src.Port()), DstPort: layers.UDPPort(dst.Port())}
+ 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 {
+ 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, 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 {
+ 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 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()}
+ setChecksum(ip)
+ return ip
+ }
+ ip := &layers.IPv6{Version: 6, HopLimit: 64, NextHeader: proto, SrcIP: src.Addr().AsSlice(), DstIP: dst.Addr().AsSlice()}
+ setChecksum(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..7c1e212
--- /dev/null
+++ b/internal/capture/session_test.go
@@ -0,0 +1,349 @@
+////----------------------------------------------------------------------------
+// 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"
+ "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 func() { _ = 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 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)
+ }
+ })
+}
+
+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)
+ }
+}
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..914f19d
--- /dev/null
+++ b/internal/dnsserver/capture_test.go
@@ -0,0 +1,79 @@
+// //----------------------------------------------------------------------------
+// // 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"
+ "net/netip"
+ "strings"
+ "testing"
+ "time"
+
+ "github.com/miekg/dns"
+
+ "github.com/lachlanharrisdev/gonetsim/internal/capture"
+ "github.com/lachlanharrisdev/gonetsim/internal/service"
+ "github.com/lachlanharrisdev/gonetsim/internal/testutil"
+)
+
+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)
+ })
+ }
+}
+
+func baseCaptureConfig(t *testing.T, network string) Config {
+ t.Helper()
+ return Config{
+ Addr: testutil.FreePort(t, "tcp"),
+ 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
+}
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..2621b4e 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)
+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},
+ }
+ client, addr, conf, 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)
+ for _, tc := range cases {
+ t.Run(tc.name, func(t *testing.T) {
+ 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)
- }
-}
-
-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)
+ if got := txt.Txt[0]; got != conf.SinkholeTXT {
+ t.Fatalf("expected %s, got %s", conf.SinkholeTXT, 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..6b59d37 100644
--- a/internal/handler/handler_test.go
+++ b/internal/handler/handler_test.go
@@ -1,58 +1,26 @@
+////----------------------------------------------------------------------------
+// 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"
- "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
+func testLogger() *slog.Logger {
+ return testutil.Logger()
}
// servePipe runs h against one end of a pipe; the other end is returned for
@@ -85,64 +53,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: testLogger()})
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) {
- reply, err := EchoHandler{}.HandleUDP(t.Context(), []byte("query"), addr, Env{Logger: discardLogger()})
+ addr, _ := net.ResolveUDPAddr("udp", "203.0.113.10:53")
+ 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)
}
})
-
- 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)
}
}
}
@@ -153,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 {
@@ -161,173 +92,60 @@ 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)
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 and log", func(t *testing.T) {
- h, err := NewLua("testdata/capture.lua", nil)
- if err != nil {
- t.Fatalf("NewLua: %v", err)
+func TestLuaState(t *testing.T) {
+ h, err := NewLua("testdata/state.lua", state.NewBudget(state.DefaultTotalLimit))
+ if err != nil {
+ t.Fatalf("NewLua: %v", err)
+ }
+ 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))
+ 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)
}
- 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)
- }
- })
-
- t.Run("sandbox globals", func(t *testing.T) {
- h, err := NewLua("testdata/sandbox_report.lua", nil)
- if err != nil {
- t.Fatalf("NewLua: %v", err)
- }
- 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
- })
+ }
}
-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 TestSandboxGlobals(t *testing.T) {
+ h, err := NewLua("testdata/sandbox_report.lua", nil)
+ 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
- })
+ client, done := servePipe(t, h, Env{Logger: testLogger()})
+ buf := make([]byte, 1024)
+ n, err := client.Read(buf)
+ if err != nil {
+ t.Fatalf("Read: %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
- }
- }
- }()
+ 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)
}
-
- 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.Close()
+ <-done
}
diff --git a/internal/handler/lua.go b/internal/handler/lua.go
index a5b1503..ed3601b 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"
)
@@ -114,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)
@@ -148,344 +144,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..61e18f3
--- /dev/null
+++ b/internal/handler/luabindings.go
@@ -0,0 +1,246 @@
+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 {
+ return pushTCPIP(L, lc.RemoteAddr())
+ }))
+
+ L.SetField(conn, "remote_port", L.NewFunction(func(L *lua.LState) int {
+ return pushTCPPort(L, lc.RemoteAddr())
+ }))
+
+ L.SetField(conn, "local_port", L.NewFunction(func(L *lua.LState) int {
+ return pushTCPPort(L, lc.LocalAddr())
+ }))
+
+ 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
+}
+
+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
new file mode 100644
index 0000000..9d98428
--- /dev/null
+++ b/internal/handler/luaconn.go
@@ -0,0 +1,124 @@
+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) {
+ return eofValue(sb.String()), 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) {
+ 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/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..c445ae4
--- /dev/null
+++ b/internal/httpserver/capture_test.go
@@ -0,0 +1,57 @@
+// //----------------------------------------------------------------------------
+// // 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"
+ "io"
+ "net/http"
+ "testing"
+ "time"
+
+ "github.com/lachlanharrisdev/gonetsim/internal/capture"
+ "github.com/lachlanharrisdev/gonetsim/internal/service"
+ "github.com/lachlanharrisdev/gonetsim/internal/testutil"
+)
+
+func TestService_CapturesHTTP(t *testing.T) {
+ conf := Config{
+ Addr: testutil.FreeTCPAddr(t),
+ StatusCode: http.StatusOK,
+ Mode: "fake",
+ Capture: true,
+ }
+ run, path := testutil.NewPcapRun(t)
+ svc, errCh := startHTTPService(t, conf, run)
+
+ get := func(url string) *http.Response {
+ _, resp := testutil.RetryGet(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()
+
+ testutil.WaitForPayloadContains(t, path, "GET /hello", 3*time.Second)
+ testutil.WaitForPayloadContains(t, path, "HTTP/1.1 200", 3*time.Second)
+
+ _ = svc.Stop(context.Background())
+ testutil.DiscardServiceStartErr(t, errCh)
+}
+
+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
+}
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/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 f8b7a7b..3f1fe44 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
@@ -160,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,
@@ -211,7 +198,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",
@@ -256,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",
@@ -376,178 +271,36 @@ 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_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) {
+func TestRealHandler_TraversalBlocked(t *testing.T) {
+ // 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)
}
-
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 {
+ 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)
- 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))
+ _ = 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))
+ }
}
}
// --- 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",
@@ -557,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 := slog.New(slog.NewTextHandler(io.Discard, nil))
- _, 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 5bc98d1..4c47729 100644
--- a/internal/httpserver/server.go
+++ b/internal/httpserver/server.go
@@ -4,12 +4,41 @@ import (
"context"
"crypto/tls"
"errors"
+ "io"
"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 +71,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) {
@@ -70,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/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..13501ab 100644
--- a/internal/listener/listener_test.go
+++ b/internal/listener/listener_test.go
@@ -1,49 +1,35 @@
+////----------------------------------------------------------------------------
+// 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 (
"context"
- "crypto/tls"
"io"
"log/slog"
"net"
"os"
- "path/filepath"
"strings"
"testing"
"time"
- "github.com/lachlanharrisdev/gonetsim/internal/capture"
+ "github.com/google/gopacket"
+ "github.com/google/gopacket/layers"
+ "github.com/google/gopacket/pcapgo"
"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 {
- 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")
- if err != nil {
- t.Fatalf("Listen: %v", err)
- }
- addr := ln.Addr().String()
- _ = ln.Close()
- return addr
-}
-
-// 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())
@@ -80,71 +66,75 @@ 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,
- 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)
- 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)
- }
- time.Sleep(20 * time.Millisecond)
- }
- t.Fatalf("no capture file appeared in %s", dir)
- return ""
+ testutil.WaitFor(t, 3*time.Second, "payload sequence match", func() bool {
+ got, err := transportPayloads(path, proto)
+ return err == nil && strings.Join(got, "|") == strings.Join(want, "|")
+ })
}
-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()
+ testutil.WaitFor(t, 3*time.Second, "payload substring match", func() bool {
+ got, err := transportPayloads(path, proto)
+ return err == nil && strings.Contains(strings.Join(got, "|"), want)
+ })
+}
+
+// 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 func() { _ = 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 := testutil.NewPcapRun(t)
- svc, err := NewService(conf, nil, testLogger())
+ svc, err := NewService(conf, nil, testLogger(), run)
if err != nil {
t.Fatalf("NewService: %v", err)
}
@@ -163,16 +153,15 @@ func TestTCPService(t *testing.T) {
}
_ = conn.Close()
- if got := captureFile(t, dir, conf.Name); got != "hello\n" {
- t.Fatalf("capture content %q", got)
- }
+ waitTransportFrames(t, path, "tcp",
+ []string{"", "", "hello\n", "hello\n", "", ""})
})
t.Run("idle timeout closes connection", func(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 +178,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: testutil.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 +212,7 @@ func TestTCPService(t *testing.T) {
})
}
-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())
- 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)
- }
- })
-
- t.Run("SNI visible to script", func(t *testing.T) {
- conf := luaConfig(t, "snitest", "sni.lua")
- conf.TLS = &tlsprovider.Config{}
-
- svc, err := NewService(conf, nil, testLogger())
- if err != nil {
- t.Fatalf("NewService: %v", err)
- }
- startService(t, svc)
-
- conn := dialTLS(t, conf.Addr, "c2.evil.example")
- defer func() { _ = conn.Close() }()
- buf := make([]byte, len("sni:c2.evil.example"))
- 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)
- }
- })
-}
-
-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,32 +243,36 @@ 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"),
+ Addr: testutil.FreePort(t, "udp"),
HandlerSpec: "builtin:echo",
- ReadTimeout: 5 * time.Second,
+ ReadTimeout: 150 * time.Millisecond,
+ Capture: true,
}
- svc, err := NewService(conf, nil, testLogger())
+ run, path := testutil.NewPcapRun(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",
- Addr: freePort(t, "udp"),
+ Addr: testutil.FreePort(t, "udp"),
HandlerSpec: "lua:packet.lua",
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,47 +281,16 @@ 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{
Name: "canceled-" + network,
Network: network,
- Addr: freePort(t, network),
+ Addr: testutil.FreePort(t, network),
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 +342,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..0369ca3 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}, nil
}
diff --git a/internal/listener/tcp.go b/internal/listener/tcp.go
index 67f3b55..274592a 100644
--- a/internal/listener/tcp.go
+++ b/internal/listener/tcp.go
@@ -10,52 +10,71 @@ 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
+ 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 +102,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.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 600821f..0000000
--- a/internal/service/manager_test.go
+++ /dev/null
@@ -1,89 +0,0 @@
-package service
-
-import (
- "context"
- "errors"
- "io"
- "log/slog"
- "strings"
- "testing"
- "time"
-)
-
-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 slog.New(slog.NewTextHandler(io.Discard, nil))
-}
-
-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 fa46b4e..18af17b 100644
--- a/internal/state/state.go
+++ b/internal/state/state.go
@@ -3,8 +3,6 @@ package state
import (
"errors"
"fmt"
- "strconv"
- "strings"
"sync"
)
@@ -50,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
}
@@ -89,38 +85,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..54551ef
--- /dev/null
+++ b/internal/testutil/testutil.go
@@ -0,0 +1,175 @@
+package testutil
+
+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 {
+ 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)
+}
+
+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 func() { _ = 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/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..c32031b 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 (
@@ -20,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)
}
@@ -69,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{
@@ -95,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)
@@ -122,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)
}
@@ -150,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")
}
}