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") } }