diff --git a/cmd/rvbox-server/main.go b/cmd/rvbox-server/main.go index 45350e9..16c3929 100644 --- a/cmd/rvbox-server/main.go +++ b/cmd/rvbox-server/main.go @@ -80,6 +80,19 @@ func run(configPath string) error { } grpcServer := grpc.NewServer(grpc.MaxRecvMsgSize(int(configured.Protocol.MaxControlRequestBytes)), grpc.MaxSendMsgSize(int(configured.Protocol.MaxControlRequestBytes))) rvboxv1.RegisterControlServer(grpcServer, controlService) + var rpcListener net.Listener + var rpcServer *http.Server + if configured.JSONRPC.Enabled { + rpcListener, err = net.Listen("tcp", configured.JSONRPC.Listen) + if err != nil { + return fmt.Errorf("listen for JSON-RPC: %w", err) + } + defer rpcListener.Close() + if configured.JSONRPC.NonLoopbackBind { + log.Printf("rvbox-server: WARNING JSON-RPC is unauthenticated and bound to non-loopback address %s", configured.JSONRPC.Listen) + } + rpcServer = &http.Server{Handler: control.NewJSONRPCHandler(controlService, int64(configured.Protocol.MaxJSONRPCBodyBytes)), ReadHeaderTimeout: configured.Flow.WriteDeadline} + } agent := &session.AgentServer{ Store: persistence, Registry: registry, Path: configured.Server.AgentPath, Limits: agentproto.Limits{ @@ -93,9 +106,12 @@ func run(configPath string) error { HeartbeatIdle: configured.Protocol.HeartbeatIdle, LivenessTimeout: configured.Protocol.LivenessTimeout, } httpServer := &http.Server{Handler: agent, ReadHeaderTimeout: configured.Flow.WriteDeadline} - serveError := make(chan error, 2) + serveError := make(chan error, 3) go func() { serveError <- httpServer.Serve(listener) }() go func() { serveError <- grpcServer.Serve(controlListener) }() + if rpcServer != nil { + go func() { serveError <- rpcServer.Serve(rpcListener) }() + } signals := make(chan os.Signal, 1) signal.Notify(signals, os.Interrupt, syscall.SIGTERM) @@ -103,15 +119,27 @@ func run(configPath string) error { select { case err := <-serveError: if errors.Is(err, http.ErrServerClosed) { + if rpcServer != nil { + _ = rpcServer.Close() + } grpcServer.Stop() return nil } + _ = httpServer.Close() + if rpcServer != nil { + _ = rpcServer.Close() + } grpcServer.Stop() return err case <-signals: shutdownContext, cancel := context.WithTimeout(context.Background(), configured.Server.ShutdownGrace) defer cancel() httpErr := httpServer.Shutdown(shutdownContext) + if rpcServer != nil { + if err := rpcServer.Shutdown(shutdownContext); httpErr == nil { + httpErr = err + } + } grpcDone := make(chan struct{}) go func() { grpcServer.GracefulStop(); close(grpcDone) }() select { diff --git a/internal/server/control/jsonrpc.go b/internal/server/control/jsonrpc.go new file mode 100644 index 0000000..131002f --- /dev/null +++ b/internal/server/control/jsonrpc.go @@ -0,0 +1,241 @@ +package control + +import ( + "bytes" + "context" + "encoding/json" + "io" + "net/http" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" + "google.golang.org/protobuf/encoding/protojson" + "google.golang.org/protobuf/proto" +) + +const ( + jsonRPCVersion = "2.0" + jsonRPCParseError = -32700 + jsonRPCInvalid = -32600 + jsonRPCMethodAbsent = -32601 + jsonRPCServerError = -32000 + defaultJSONRPCBody = 24 << 20 +) + +// JSONRPCHandler adapts unary Control methods to standard JSON-RPC 2.0. It +// intentionally has no authentication in v1; callers should keep it loopback +// bound unless they explicitly accept the debugging exposure. +type JSONRPCHandler struct { + Service *Service + MaxBody int64 + Decoder protojson.UnmarshalOptions + Encoder protojson.MarshalOptions +} + +func NewJSONRPCHandler(service *Service, maxBody int64) *JSONRPCHandler { + if maxBody <= 0 { + maxBody = defaultJSONRPCBody + } + return &JSONRPCHandler{Service: service, MaxBody: maxBody, Decoder: protojson.UnmarshalOptions{DiscardUnknown: false}, Encoder: protojson.MarshalOptions{UseProtoNames: false}} +} + +func (handler *JSONRPCHandler) ServeHTTP(response http.ResponseWriter, request *http.Request) { + if request.Method != http.MethodPost { + response.Header().Set("Allow", http.MethodPost) + http.Error(response, "method not allowed", http.StatusMethodNotAllowed) + return + } + if handler == nil || handler.Service == nil { + http.Error(response, "control service unavailable", http.StatusServiceUnavailable) + return + } + body, err := io.ReadAll(io.LimitReader(request.Body, handler.MaxBody+1)) + if err != nil { + handler.writeRPCError(response, nil, jsonRPCParseError, "could not read request", nil) + return + } + if int64(len(body)) > handler.MaxBody { + handler.writeRPCError(response, nil, jsonRPCInvalid, "request body exceeds limit", nil) + return + } + var envelope jsonRPCRequest + decoder := json.NewDecoder(bytes.NewReader(body)) + decoder.DisallowUnknownFields() + if err := decoder.Decode(&envelope); err != nil { + handler.writeRPCError(response, nil, jsonRPCParseError, "invalid JSON", nil) + return + } + var trailing any + if err := decoder.Decode(&trailing); err != io.EOF || envelope.JSONRPC != jsonRPCVersion || envelope.Method == "" || len(envelope.ID) == 0 || bytes.Equal(bytes.TrimSpace(envelope.ID), []byte("null")) { + handler.writeRPCError(response, nil, jsonRPCInvalid, "invalid JSON-RPC request", nil) + return + } + result, callErr := handler.call(request.Context(), envelope.Method, envelope.Params) + if callErr != nil { + code, message, data := jsonRPCDomainError(callErr) + handler.writeRPCError(response, envelope.ID, code, message, data) + return + } + handler.writeRPCResult(response, envelope.ID, result) +} + +type jsonRPCRequest struct { + JSONRPC string `json:"jsonrpc"` + ID json.RawMessage `json:"id"` + Method string `json:"method"` + Params json.RawMessage `json:"params"` +} + +type jsonRPCResponse struct { + JSONRPC string `json:"jsonrpc"` + ID json.RawMessage `json:"id"` + Result json.RawMessage `json:"result,omitempty"` + Error *jsonRPCError `json:"error,omitempty"` +} + +type jsonRPCError struct { + Code int `json:"code"` + Message string `json:"message"` + Data json.RawMessage `json:"data,omitempty"` +} + +func (handler *JSONRPCHandler) call(ctx context.Context, method string, params json.RawMessage) (proto.Message, error) { + decode := func(message proto.Message) error { + if len(params) == 0 || bytes.Equal(bytes.TrimSpace(params), []byte("null")) { + params = []byte("{}") + } + return handler.Decoder.Unmarshal(params, message) + } + switch method { + case "listClients": + input := &rvboxv1.ListClientsRequest{} + if err := decode(input); err != nil { + return nil, status.Error(codes.InvalidArgument, err.Error()) + } + return handler.Service.ListClients(ctx, input) + case "getClient": + input := &rvboxv1.GetClientRequest{} + if err := decode(input); err != nil { + return nil, status.Error(codes.InvalidArgument, err.Error()) + } + return handler.Service.GetClient(ctx, input) + case "listCommands": + input := &rvboxv1.ListCommandsRequest{} + if err := decode(input); err != nil { + return nil, status.Error(codes.InvalidArgument, err.Error()) + } + return handler.Service.ListCommands(ctx, input) + case "getCommand": + input := &rvboxv1.GetCommandRequest{} + if err := decode(input); err != nil { + return nil, status.Error(codes.InvalidArgument, err.Error()) + } + return handler.Service.GetCommand(ctx, input) + case "runCommand": + input := &rvboxv1.RunCommandRequest{} + if err := decode(input); err != nil { + return nil, status.Error(codes.InvalidArgument, err.Error()) + } + return handler.Service.RunCommand(ctx, input) + case "appendStdin": + input := &rvboxv1.AppendStdinRequest{} + if err := decode(input); err != nil { + return nil, status.Error(codes.InvalidArgument, err.Error()) + } + return handler.Service.AppendStdin(ctx, input) + case "closeStdin": + input := &rvboxv1.CloseStdinRequest{} + if err := decode(input); err != nil { + return nil, status.Error(codes.InvalidArgument, err.Error()) + } + return handler.Service.CloseStdin(ctx, input) + case "signalCommand": + input := &rvboxv1.ControlSignalCommandRequest{} + if err := decode(input); err != nil { + return nil, status.Error(codes.InvalidArgument, err.Error()) + } + return handler.Service.SignalCommand(ctx, input) + case "getOutput": + input := &rvboxv1.GetOutputRequest{} + if err := decode(input); err != nil { + return nil, status.Error(codes.InvalidArgument, err.Error()) + } + return handler.Service.GetOutput(ctx, input) + case "listStorageIncidents": + input := &rvboxv1.ListStorageIncidentsRequest{} + if err := decode(input); err != nil { + return nil, status.Error(codes.InvalidArgument, err.Error()) + } + return handler.Service.ListStorageIncidents(ctx, input) + case "repairStorageIncident": + input := &rvboxv1.RepairStorageIncidentRequest{} + if err := decode(input); err != nil { + return nil, status.Error(codes.InvalidArgument, err.Error()) + } + return handler.Service.RepairStorageIncident(ctx, input) + case "acknowledgeStorageIncident": + input := &rvboxv1.AcknowledgeStorageIncidentRequest{} + if err := decode(input); err != nil { + return nil, status.Error(codes.InvalidArgument, err.Error()) + } + return handler.Service.AcknowledgeStorageIncident(ctx, input) + case "authorizeClientTakeover": + input := &rvboxv1.AuthorizeClientTakeoverRequest{} + if err := decode(input); err != nil { + return nil, status.Error(codes.InvalidArgument, err.Error()) + } + return handler.Service.AuthorizeClientTakeover(ctx, input) + default: + return nil, status.Error(codes.Unimplemented, "method not found") + } +} + +func (handler *JSONRPCHandler) writeRPCResult(response http.ResponseWriter, id json.RawMessage, message proto.Message) { + encoded, err := handler.Encoder.Marshal(message) + if err != nil { + handler.writeRPCError(response, id, jsonRPCServerError, "could not encode response", nil) + return + } + response.Header().Set("Content-Type", "application/json") + response.WriteHeader(http.StatusOK) + _, _ = response.Write(mustJSON(jsonRPCResponse{JSONRPC: jsonRPCVersion, ID: append([]byte(nil), id...), Result: encoded})) +} + +func (handler *JSONRPCHandler) writeRPCError(response http.ResponseWriter, id json.RawMessage, code int, message string, data json.RawMessage) { + if len(id) == 0 { + id = []byte("null") + } + response.Header().Set("Content-Type", "application/json") + response.WriteHeader(http.StatusOK) + _, _ = response.Write(mustJSON(jsonRPCResponse{JSONRPC: jsonRPCVersion, ID: append([]byte(nil), id...), Error: &jsonRPCError{Code: code, Message: message, Data: data}})) +} + +func jsonRPCDomainError(err error) (int, string, json.RawMessage) { + if err == nil { + return 0, "", nil + } + converted := status.Convert(err) + if converted.Code() == codes.Unimplemented && converted.Message() == "method not found" { + return jsonRPCMethodAbsent, converted.Message(), nil + } + data := []byte(nil) + for _, detail := range converted.Details() { + if control, ok := detail.(*rvboxv1.ControlError); ok { + data, _ = protojson.Marshal(control) + break + } + } + if converted.Code() == codes.InvalidArgument { + return jsonRPCInvalid, converted.Message(), data + } + return jsonRPCServerError, converted.Message(), data +} + +func mustJSON(value any) []byte { + encoded, err := json.Marshal(value) + if err != nil { + return []byte(`{"jsonrpc":"2.0","id":null,"error":{"code":-32000,"message":"internal encoding error"}}`) + } + return encoded +} diff --git a/internal/server/control/jsonrpc_test.go b/internal/server/control/jsonrpc_test.go new file mode 100644 index 0000000..8176d61 --- /dev/null +++ b/internal/server/control/jsonrpc_test.go @@ -0,0 +1,97 @@ +package control + +import ( + "bytes" + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "path/filepath" + "testing" + "time" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/server/store" + "google.golang.org/protobuf/encoding/protojson" +) + +func TestJSONRPCUnaryMatchesControlService_HP_CTL_15(t *testing.T) { + persistence, err := store.Open(context.Background(), store.Options{DataDir: filepath.Join(t.TempDir(), "state"), BusyTimeout: time.Second}) + if err != nil { + t.Fatal(err) + } + defer persistence.Close() + registerControlClient(t, persistence, "win-a", rvboxv1.Platform_PLATFORM_WINDOWS, rvboxv1.ShellType_SHELL_POWERSHELL, 9) + service, err := NewService(Options{Store: persistence, CursorKey: bytesKey()}) + if err != nil { + t.Fatal(err) + } + handler := httptest.NewServer(NewJSONRPCHandler(service, 1<<20)) + defer handler.Close() + requestID := fixedIssue(0xaf).String() + params, err := protojson.Marshal(&rvboxv1.RunCommandRequest{TargetClientId: "win-a", RequestId: requestID, Spec: &rvboxv1.ExecutionSpec{Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "echo rpc"}}}) + if err != nil { + t.Fatal(err) + } + body, _ := json.Marshal(map[string]any{"jsonrpc": "2.0", "id": 7, "method": "runCommand", "params": json.RawMessage(params)}) + response, err := http.Post(handler.URL, "application/json", bytes.NewReader(body)) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + encoded, _ := io.ReadAll(response.Body) + var envelope struct { + Result json.RawMessage `json:"result"` + Error json.RawMessage `json:"error"` + } + if err := json.Unmarshal(encoded, &envelope); err != nil || len(envelope.Error) != 0 { + t.Fatalf("JSON-RPC response = %s, %v", encoded, err) + } + result := &rvboxv1.RunCommandResponse{} + if err := protojson.Unmarshal(envelope.Result, result); err != nil || result.GetIssueUuid() != requestID { + t.Fatalf("JSON-RPC result = %s, %v", envelope.Result, err) + } + unknown, _ := json.Marshal(map[string]any{"jsonrpc": "2.0", "id": "x", "method": "notARealMethod", "params": map[string]any{}}) + unknownResponse, err := http.Post(handler.URL, "application/json", bytes.NewReader(unknown)) + if err != nil { + t.Fatal(err) + } + defer unknownResponse.Body.Close() + unknownBody, _ := io.ReadAll(unknownResponse.Body) + if !bytes.Contains(unknownBody, []byte(`-32601`)) { + t.Fatalf("unknown method response = %s", unknownBody) + } +} + +func TestJSONRPCRejectsOversizeAndMalformedRequests_BH_CTL_16(t *testing.T) { + persistence, err := store.Open(context.Background(), store.Options{DataDir: filepath.Join(t.TempDir(), "state"), BusyTimeout: time.Second}) + if err != nil { + t.Fatal(err) + } + defer persistence.Close() + service, err := NewService(Options{Store: persistence, CursorKey: bytesKey()}) + if err != nil { + t.Fatal(err) + } + handler := httptest.NewServer(NewJSONRPCHandler(service, 16)) + defer handler.Close() + response, err := http.Post(handler.URL, "application/json", bytes.NewReader([]byte(`{"jsonrpc":"2.0","id":1,"method":"listClients","params":{}}`))) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + body, _ := io.ReadAll(response.Body) + if !bytes.Contains(body, []byte(`-32600`)) { + t.Fatalf("oversize response = %s", body) + } + malformed, err := http.Post(handler.URL, "application/json", bytes.NewReader([]byte(`{"jsonrpc":"2.0","id":1,"method":"listClients","params":{}}`[:16]))) + if err != nil { + t.Fatal(err) + } + defer malformed.Body.Close() + malformedBody, _ := io.ReadAll(malformed.Body) + if !bytes.Contains(malformedBody, []byte(`-32700`)) { + t.Fatalf("malformed response = %s", malformedBody) + } +} diff --git a/test/coverage.toml b/test/coverage.toml index a41c452..25106ac 100644 --- a/test/coverage.toml +++ b/test/coverage.toml @@ -89,6 +89,18 @@ layer = "unit" status = "implemented" tests = ["internal/server/control/service_test.go:TestStorageIncidentControlLifecycle_HP_CONTROL_14"] +[[requirements]] +id = "HP-CTL-15" +layer = "integration" +status = "implemented" +tests = ["internal/server/control/jsonrpc_test.go:TestJSONRPCUnaryMatchesControlService_HP_CTL_15"] + +[[requirements]] +id = "BH-CTL-16" +layer = "integration" +status = "implemented" +tests = ["internal/server/control/jsonrpc_test.go:TestJSONRPCRejectsOversizeAndMalformedRequests_BH_CTL_16"] + [[requirements]] id = "BH-CTL-01" layer = "unit"