167 lines
3.9 KiB
Go
167 lines
3.9 KiB
Go
package socket_test
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"path/filepath"
|
|
"sync"
|
|
"testing"
|
|
"time"
|
|
|
|
"gitea.alexandru.macocian.me/Charlie/sherlock/internal/rpc"
|
|
"gitea.alexandru.macocian.me/Charlie/sherlock/internal/socket"
|
|
)
|
|
|
|
func newServer(t *testing.T, handler socket.Handler) (string, func()) {
|
|
t.Helper()
|
|
path := filepath.Join(t.TempDir(), "test.sock")
|
|
srv, err := socket.Listen(path, handler)
|
|
if err != nil {
|
|
t.Fatalf("Listen: %v", err)
|
|
}
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
done := make(chan struct{})
|
|
go func() {
|
|
defer close(done)
|
|
if err := srv.Serve(ctx); err != nil {
|
|
t.Errorf("Serve: %v", err)
|
|
}
|
|
}()
|
|
return path, func() {
|
|
cancel()
|
|
_ = srv.Close()
|
|
<-done
|
|
}
|
|
}
|
|
|
|
func TestRoundTrip_Success(t *testing.T) {
|
|
handler := socket.HandlerFunc(func(_ context.Context, req rpc.Request) rpc.Response {
|
|
if req.Method != "echo" {
|
|
return rpc.Response{Error: rpc.NewError(rpc.CodeUnknownMethod, "nope")}
|
|
}
|
|
out, _ := rpc.MarshalResult(map[string]string{"got": "ok"})
|
|
return rpc.Response{Result: out}
|
|
})
|
|
path, stop := newServer(t, handler)
|
|
defer stop()
|
|
|
|
c, err := socket.Dial(path)
|
|
if err != nil {
|
|
t.Fatalf("Dial: %v", err)
|
|
}
|
|
defer func() { _ = c.Close() }()
|
|
|
|
var resp map[string]string
|
|
if err := c.Call(context.Background(), "echo", nil, &resp); err != nil {
|
|
t.Fatalf("Call: %v", err)
|
|
}
|
|
if resp["got"] != "ok" {
|
|
t.Fatalf("got %v", resp)
|
|
}
|
|
}
|
|
|
|
func TestRoundTrip_Error(t *testing.T) {
|
|
handler := socket.HandlerFunc(func(_ context.Context, req rpc.Request) rpc.Response {
|
|
return rpc.Response{Error: rpc.NewError(rpc.CodeNotLoggedIn, "session missing")}
|
|
})
|
|
path, stop := newServer(t, handler)
|
|
defer stop()
|
|
|
|
c, err := socket.Dial(path)
|
|
if err != nil {
|
|
t.Fatalf("Dial: %v", err)
|
|
}
|
|
defer func() { _ = c.Close() }()
|
|
|
|
err = c.Call(context.Background(), "status", nil, nil)
|
|
if err == nil {
|
|
t.Fatal("expected error")
|
|
}
|
|
rpcErr, ok := err.(*rpc.Error)
|
|
if !ok {
|
|
t.Fatalf("expected *rpc.Error, got %T", err)
|
|
}
|
|
if rpcErr.Code != rpc.CodeNotLoggedIn {
|
|
t.Fatalf("code = %v", rpcErr.Code)
|
|
}
|
|
}
|
|
|
|
func TestRoundTrip_BadJSON(t *testing.T) {
|
|
handler := socket.HandlerFunc(func(_ context.Context, req rpc.Request) rpc.Response {
|
|
return rpc.Response{}
|
|
})
|
|
path, stop := newServer(t, handler)
|
|
defer stop()
|
|
|
|
// Send garbage directly so we exercise the server-side decode path.
|
|
conn, err := dialRaw(path)
|
|
if err != nil {
|
|
t.Fatalf("dialRaw: %v", err)
|
|
}
|
|
defer func() { _ = conn.Close() }()
|
|
|
|
if _, err := conn.Write([]byte("not json\n")); err != nil {
|
|
t.Fatalf("Write: %v", err)
|
|
}
|
|
buf := make([]byte, 1024)
|
|
n, err := conn.Read(buf)
|
|
if err != nil {
|
|
t.Fatalf("Read: %v", err)
|
|
}
|
|
var resp rpc.Response
|
|
if err := json.Unmarshal(buf[:n], &resp); err != nil {
|
|
t.Fatalf("Unmarshal: %v", err)
|
|
}
|
|
if resp.Error == nil || resp.Error.Code != rpc.CodeBadRequest {
|
|
t.Fatalf("expected bad_request, got %+v", resp.Error)
|
|
}
|
|
}
|
|
|
|
func TestConcurrentCalls(t *testing.T) {
|
|
handler := socket.HandlerFunc(func(_ context.Context, req rpc.Request) rpc.Response {
|
|
var in struct{ N int }
|
|
_ = json.Unmarshal(req.Params, &in)
|
|
out, _ := rpc.MarshalResult(map[string]int{"n": in.N})
|
|
return rpc.Response{Result: out}
|
|
})
|
|
path, stop := newServer(t, handler)
|
|
defer stop()
|
|
|
|
const goroutines = 20
|
|
const callsEach = 50
|
|
|
|
var wg sync.WaitGroup
|
|
for i := 0; i < goroutines; i++ {
|
|
wg.Add(1)
|
|
go func() {
|
|
defer wg.Done()
|
|
c, err := socket.Dial(path)
|
|
if err != nil {
|
|
t.Errorf("Dial: %v", err)
|
|
return
|
|
}
|
|
defer func() { _ = c.Close() }()
|
|
for j := 0; j < callsEach; j++ {
|
|
var out map[string]int
|
|
err := c.Call(context.Background(), "echo",
|
|
map[string]int{"N": j}, &out)
|
|
if err != nil {
|
|
t.Errorf("Call: %v", err)
|
|
return
|
|
}
|
|
if out["n"] != j {
|
|
t.Errorf("got %d want %d", out["n"], j)
|
|
return
|
|
}
|
|
}
|
|
}()
|
|
}
|
|
doneCh := make(chan struct{})
|
|
go func() { wg.Wait(); close(doneCh) }()
|
|
select {
|
|
case <-doneCh:
|
|
case <-time.After(10 * time.Second):
|
|
t.Fatal("concurrent calls timed out")
|
|
}
|
|
}
|