91 lines
2.5 KiB
Go
91 lines
2.5 KiB
Go
package identity_test
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"os"
|
|
"path/filepath"
|
|
"testing"
|
|
|
|
"git.wntrmute.dev/kyle/crossbar/internal/identity"
|
|
)
|
|
|
|
func TestParseWhois(t *testing.T) {
|
|
raw, err := os.ReadFile(filepath.Join("testdata", "whois.json"))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
id, err := identity.ParseWhois(raw)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if id.Node != "titan" || id.Login == "" {
|
|
t.Errorf("parsed %+v, want Node titan and a login", id)
|
|
}
|
|
if _, err := identity.ParseWhois([]byte(`{"Node":{}}`)); err == nil {
|
|
t.Error("a whois answer without a node name must be an error")
|
|
}
|
|
if _, err := identity.ParseWhois([]byte(`nope`)); err == nil {
|
|
t.Error("non-JSON must be an error")
|
|
}
|
|
}
|
|
|
|
// fakeResolver answers from a map; "" means not a tailnet peer.
|
|
type fakeResolver map[string]string
|
|
|
|
func (f fakeResolver) Identity(ctx context.Context, ip string) (identity.ID, error) {
|
|
n, ok := f[ip]
|
|
if !ok {
|
|
return identity.ID{}, identity.ErrNotAPeer
|
|
}
|
|
return identity.ID{Node: n, Login: n + "@example"}, nil
|
|
}
|
|
|
|
func TestChecker(t *testing.T) {
|
|
c := identity.NewChecker(fakeResolver{"100.64.0.5": "talos", "100.64.0.9": "titan"})
|
|
for _, tc := range []struct {
|
|
name string
|
|
peers []string
|
|
addr string
|
|
want error
|
|
}{
|
|
{"open route", nil, "203.0.113.7:1", nil},
|
|
{"allowed peer", []string{"talos", "titan"}, "100.64.0.5:44444", nil},
|
|
{"other peer", []string{"talos"}, "100.64.0.9:1", identity.ErrForbidden},
|
|
{"not a peer", []string{"talos"}, "203.0.113.7:1", identity.ErrForbidden},
|
|
{"loopback", []string{"talos"}, "127.0.0.1:1", identity.ErrForbidden},
|
|
{"garbage addr", []string{"talos"}, "nonsense", identity.ErrForbidden},
|
|
} {
|
|
t.Run(tc.name, func(t *testing.T) {
|
|
got := c.Allow(context.Background(), tc.peers, tc.addr)
|
|
if !errors.Is(got, tc.want) && !(got == nil && tc.want == nil) {
|
|
t.Errorf("Allow(%v, %q) = %v, want %v", tc.peers, tc.addr, got, tc.want)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestCheckerCachesPerAddress(t *testing.T) {
|
|
calls := 0
|
|
r := countingResolver{f: fakeResolver{"100.64.0.5": "talos"}, calls: &calls}
|
|
c := identity.NewChecker(r)
|
|
for i := 0; i < 5; i++ {
|
|
if err := c.Allow(context.Background(), []string{"talos"}, "100.64.0.5:1"); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
if calls != 1 {
|
|
t.Errorf("resolver called %d times for one address, want 1 (cache)", calls)
|
|
}
|
|
}
|
|
|
|
type countingResolver struct {
|
|
f fakeResolver
|
|
calls *int
|
|
}
|
|
|
|
func (c countingResolver) Identity(ctx context.Context, ip string) (identity.ID, error) {
|
|
*c.calls++
|
|
return c.f.Identity(ctx, ip)
|
|
}
|